"""Bounded one-model diagnosis of BF16 batching and shared-cache parity.""" import argparse from contextlib import nullcontext import json from pathlib import Path import sys parser = argparse.ArgumentParser() parser.add_argument("--run", required=True) args = parser.parse_args() run = Path(args.run) Path("/proc/self/oom_score_adj").write_text("0") sys.path.insert(0, str(Path(__file__).resolve().parents[1])) import torch 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) scorer = DecisionScorer(str(Path.home() / "ai/models/opensysone/Qwen2.5-0.5B-060db649"), dtype=torch.bfloat16).eval() saved = torch.load(run / "checkpoint.pt", map_location="cpu", weights_only=False) scorer.load_state_dict(saved["trainable_state"], strict=False) group = [json.loads(line) for line in (run / "test.jsonl").read_text().splitlines()] group = [row for row in group if row["group"] == "TEST-000"] def values(scores): return [score.float().tolist() for score in scores] def difference(first, second): return { "logit_max_abs": max((a.float() - b.float()).abs().max().item() for a, b in zip(first, second)), "probability_max_abs": max((a.float().softmax(0) - b.float().softmax(0)).abs().max().item() for a, b in zip(first, second)), } results = {"checkpoint": str(run / "checkpoint.pt"), "head_norm": scorer.head.weight.norm().item(), "modes": {}} for precision in ("bf16", "fp32"): if precision == "fp32": scorer.lm.float() for backend in ("default", "math", "flash"): key = precision + "_" + backend context = nullcontext() if backend == "default" else sdpa_kernel( SDPBackend.MATH if backend == "math" else SDPBackend.FLASH_ATTENTION) try: with torch.inference_mode(), context: full = scorer.score_examples(group) shared = scorer.scores_shared(group[0]["state"], group, 4) shared_single_chunk = scorer.scores_shared(group[0]["state"], group, 16) single = [scorer.score_examples([row])[0] for row in group] isolated = [scorer.scores_shared(row["state"], [row], 4)[0] for row in group] reversed_rows = [{**row, "choices": list(reversed(row["choices"]))} for row in group] reversed_full = [score.flip(0) for score in scorer.score_examples(reversed_rows)] reversed_shared = [score.flip(0) for score in scorer.scores_shared(group[0]["state"], reversed_rows, 4)] results["modes"][key] = { "full_vs_shared": difference(full, shared), "full_vs_single": difference(full, single), "shared_chunk_sizes": difference(shared, shared_single_chunk), "shared_vs_isolated": difference(shared, isolated), "full_permutation": difference(full, reversed_full), "shared_permutation": difference(shared, reversed_shared), "full_logits": values(full), "shared_logits": values(shared), } except Exception as error: results["modes"][key] = {"error": repr(error)} print(json.dumps({key: results["modes"][key]}), flush=True) (run / "parity_diagnosis.json").write_text(json.dumps(results, indent=2) + "\n")