"""Read-only checkpoint ablations; fresh output, bounded memory, no training. Compare all synthetic test groups, then trace the first group's final candidate token through each decoder layer. FP32 linear arithmetic with BF16 outputs is a diagnostic intervention, not a proposed training or deployment configuration. """ import argparse from collections import defaultdict from contextlib import nullcontext import hashlib import importlib.metadata import json import os from pathlib import Path import platform import subprocess import sys import time import types def sha256(path): return hashlib.sha256(Path(path).read_bytes()).hexdigest() def write_json(path, value): temporary = path.with_suffix(".tmp") temporary.write_text(json.dumps(value, indent=2) + "\n") temporary.replace(path) def available(): return int(next(line.split()[1] for line in Path("/proc/meminfo").read_text().splitlines() if line.startswith("MemAvailable:"))) * 1024 def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--run", required=True, help="Trusted project checkpoint directory") parser.add_argument("--output", required=True, help="New artifact directory") args = parser.parse_args() run, out = Path(args.run).resolve(), Path(args.output).resolve() out.mkdir(parents=True, exist_ok=False) Path("/proc/self/oom_score_adj").write_text("0") if available() < 24 * 2**30: raise RuntimeError("Requires 24 GiB available unified memory") sys.path.insert(0, str(Path(__file__).resolve().parents[1])) import torch from torch.nn import functional as F from torch.nn.attention import SDPBackend, sdpa_kernel from decision_model import DecisionScorer torch.set_num_threads(8) torch.cuda.set_per_process_memory_fraction( 16 * 2**30 / torch.cuda.get_device_properties(0).total_memory) parent = json.loads((run / "manifest.json").read_text()) model = Path(parent["config"]["model"]) provenance = json.loads((model / "opensysone-provenance.json").read_text()) if provenance != {"model_id": parent["model_id"], "revision": parent["model_revision"]}: raise ValueError("Base model provenance mismatch") manifest = { "config": vars(args), "pid": os.getpid(), "hostname": platform.node(), "started_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "git_commit": subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), "git_status": subprocess.check_output(["git", "status", "--porcelain"], text=True), "source_sha256": {str(p): sha256(p) for p in [Path(__file__), Path("decision_model.py")]}, "checkpoint_sha256": sha256(run / "checkpoint.pt"), "test_sha256": sha256(run / "test.jsonl"), "model": provenance, "packages": {p: importlib.metadata.version(p) for p in ["torch", "transformers"]}, "cuda": torch.version.cuda, "gpu": torch.cuda.get_device_name(), "initial_mem_available_bytes": available(), "cuda_cap_bytes": 16 * 2**30, "oom_score_adj": Path("/proc/self/oom_score_adj").read_text().strip(), } write_json(out / "manifest.json", manifest) scorer = DecisionScorer(str(model), dtype=torch.bfloat16).eval() saved = torch.load(run / "checkpoint.pt", map_location="cpu", weights_only=False) if saved["model_revision"] != provenance["revision"]: raise ValueError("Checkpoint base revision mismatch") scorer.load_state_dict(saved["trainable_state"], strict=False) del saved groups = defaultdict(list) for line in (run / "test.jsonl").read_text().splitlines(): row = json.loads(line) groups[row["group"]].append(row) def difference(a, b): return max((x.float().softmax(0) - y.float().softmax(0)).abs().max().item() for x, y in zip(a, b)) def trace(rows): sequences = [scorer.prefix_ids(row["state"]) + scorer.branch_ids(row["question"], choice) for row in rows for choice in row["choices"]] for sequence in sequences: if len(sequence) > scorer.max_tokens: raise ValueError("Diagnostic input exceeds scorer limit") ids, mask, lengths = scorer._pad(sequences) stages = {} handles = [] def hook(name): def capture(module, inputs, output): hidden = output[0] if isinstance(output, tuple) else output stages[name] = scorer._last(hidden, lengths).float().cpu() return capture modules = [("embedding", scorer.lm.model.embed_tokens)] first_layer = scorer.lm.model.layers[0] modules += [("layer_00_" + name, module) for name, module in first_layer.named_modules() if name and isinstance(module, (torch.nn.Linear,))] modules += [("layer_00_input_norm", first_layer.input_layernorm), ("layer_00_post_attention_norm", first_layer.post_attention_layernorm)] modules += [(f"layer_{i:02d}", layer) for i, layer in enumerate(scorer.lm.model.layers)] modules += [("final_norm", scorer.lm.model.norm)] try: for name, module in modules: handles.append(module.register_forward_hook(hook(name))) scorer.lm.model(input_ids=ids, attention_mask=mask, use_cache=False) finally: for handle in handles: handle.remove() return stages results = {"claim_scope": "Synthetic correctness diagnosis only", "modes": {}} original_forwards = {} def linear_fp32(module, value): return F.linear(value.float(), module.weight.float(), module.bias.float() if module.bias is not None else None).to(value.dtype) modes = [ ("bf16_default", False, False, False), ("bf16_strict_reduction", True, False, False), ("bf16_math_strict_reduction", True, True, False), ("bf16_fp32_linear", True, False, True), ("bf16_math_fp32_linear", True, True, True), ("fp32_reference", True, False, False), ] with torch.inference_mode(): for name, strict, math, linear in modes: if available() < 24 * 2**30: raise RuntimeError("Available unified memory fell below 24 GiB") for module, forward in original_forwards.items(): module.forward = forward original_forwards.clear() if name == "fp32_reference": scorer.lm.float() if linear: for module in scorer.lm.model.modules(): if isinstance(module, torch.nn.Linear): original_forwards[module] = module.forward module.forward = types.MethodType(linear_fp32, module) torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = not strict context = sdpa_kernel(SDPBackend.MATH) if math else nullcontext() tick = time.perf_counter() rows_result = [] with context: for group_id, rows in groups.items(): full = scorer.score_examples(rows) full_repeat = scorer.score_examples(rows) single = [scorer.score_examples([row])[0] for row in rows] shared = scorer.scores_shared(rows[0]["state"], rows, 4) shared_repeat = scorer.scores_shared(rows[0]["state"], rows, 4) wide = scorer.scores_shared(rows[0]["state"], rows, 16) isolated = [scorer.scores_shared(row["state"], [row], 4)[0] for row in rows] reversed_rows = [{**row, "choices": list(reversed(row["choices"]))} for row in rows] reversed_full = [s.flip(0) for s in scorer.score_examples(reversed_rows)] reversed_scores = [s.flip(0) for s in scorer.scores_shared(rows[0]["state"], reversed_rows, 4)] checks = {"full_vs_single": difference(full, single), "full_repeat": difference(full, full_repeat), "shared_repeat": difference(shared, shared_repeat), "full_permutation": difference(full, reversed_full), "full_vs_shared": difference(full, shared), "shared_chunk_sizes": difference(shared, wide), "shared_vs_isolated": difference(shared, isolated), "shared_permutation": difference(shared, reversed_scores)} rows_result.append({"group": group_id, "checks": checks, "full_logits": [s.float().tolist() for s in full]}) first = next(iter(groups.values())) batched = trace(first) separate = [trace([row]) for row in first] layer_trace = {} for stage, value in batched.items(): other = torch.cat([item[stage] for item in separate]) delta = value - other layer_trace[stage] = {"max_abs": delta.abs().max().item(), "rms": delta.square().mean().sqrt().item()} maxima = {key: max(row["checks"][key] for row in rows_result) for key in rows_result[0]["checks"]} results["modes"][name] = {"groups": rows_result, "worst_probability_max_abs": maxima, "groups_above_original_0_02_gate": sum( max(row["checks"].values()) > 0.02 for row in rows_result), "first_group_layer_trace": layer_trace, "seconds": time.perf_counter() - tick} write_json(out / "precision.json", results) print(json.dumps({"mode": name, "worst_probability_max_abs": maxima, "seconds": time.perf_counter() - tick}), flush=True) results["resources"] = {"peak_cuda_allocated_bytes": torch.cuda.max_memory_allocated(), "peak_cuda_reserved_bytes": torch.cuda.max_memory_reserved(), "final_mem_available_bytes": available()} results["status"] = "complete" write_json(out / "precision.json", results) if __name__ == "__main__": main()