opensysone / source /scripts /investigate_precision.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
2d5c26a verified
Raw History Blame
10.5 kB
"""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()