""" selfcheck.py -- Step 2 self-checks that need no GPU/model. Runs: 1. eigh == Cholesky whitening equivalence (matched precision, fp64). 2. lm_head hard-defense: a fake model whose head was replaced MUST raise. 3. McNemar vs scipy (delegates to analysis/mcnemar.self_test). 4. per-item schema contract: 3 mock records are unique, complete, mcnemar-readable. The 3-item REAL-eval schema check and the real-model head guard need a GPU and are run separately on the compute node. """ import os import sys import json import tempfile import torch import torch.nn as nn HERE = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, HERE) import common as C from compress.compress import whiten_truncate from analysis import mcnemar results = {} # ── 1. eigh == Cholesky (fp64 matched precision) ────────────────────────────── def check_eigh_cholesky(): torch.manual_seed(0) worst = 0.0 worst_ratio = 0.0 for (d_out, d_in) in [(128, 96), (96, 128), (256, 256)]: W = torch.randn(d_out, d_in, dtype=torch.float64) Xc = torch.randn(4000, d_in, dtype=torch.float64) XtX = (Xc.t() @ Xc) / 4000 + 1e-3 * torch.eye(d_in, dtype=torch.float64) for ratio in (0.3, 0.5, 0.7): Ac, Bc, k = whiten_truncate(W, XtX, ratio, decomp="cholesky", solve_dtype=torch.float64, out_dtype=torch.float64) Ae, Be, _ = whiten_truncate(W, XtX, ratio, decomp="eigh", solve_dtype=torch.float64, out_dtype=torch.float64) Wc = Ac @ Bc We = Ae @ Be diff = float((Wc - We).abs().max()) trunc_err = float((W - Wc).norm() / W.norm()) rel = diff / max(trunc_err, 1e-30) worst = max(worst, diff) worst_ratio = max(worst_ratio, rel) fp32_floor = float(torch.finfo(torch.float32).eps) # ~1.19e-7 ok = worst <= 2 * fp32_floor and worst_ratio < 0.01 results["eigh_cholesky"] = { "worst_abs_diff_fp64": worst, "worst_diff_over_truncerr": worst_ratio, "fp32_floor": fp32_floor, "criterion": "diff <= 2*fp32_floor AND diff/truncerr < 1%", "pass": ok, } print(f"[1] eigh==Cholesky: max|W'_chol - W'_eigh|(fp64)={worst:.3e} " f"(expect ~1e-12), diff/truncerr={worst_ratio:.3e} -> {'PASS' if ok else 'FAIL'}") return ok # ── 2. lm_head hard defense ─────────────────────────────────────────────────── class _FakeModel: """Minimal stand-in exposing get_output_embeddings() like LLaDAModelLM.""" def __init__(self, head): self._head = head def get_output_embeddings(self): return self._head def check_lm_head_guard(): # dense head -> passes ok_dense = False try: C.assert_head_dense(_FakeModel(nn.Linear(8, 16))) ok_dense = True except Exception as e: print(f" unexpected: dense head raised {e}") # replaced head -> MUST raise A = torch.randn(16, 4) B = torch.randn(4, 8) lr = C.LowRankLinear(A, B) raised = False try: C.assert_head_dense(_FakeModel(lr)) except RuntimeError: raised = True # Identity head -> MUST raise raised_id = False try: C.assert_head_dense(_FakeModel(nn.Identity())) except RuntimeError: raised_id = True ok = ok_dense and raised and raised_id results["lm_head_guard"] = {"dense_passes": ok_dense, "lowrank_raises": raised, "identity_raises": raised_id, "pass": ok} print(f"[2] lm_head guard: dense_passes={ok_dense} lowrank_raises={raised} " f"identity_raises={raised_id} -> {'PASS' if ok else 'FAIL'}") return ok # ── 3. McNemar vs scipy ─────────────────────────────────────────────────────── def check_mcnemar(): ok = mcnemar.self_test() results["mcnemar_vs_scipy"] = {"pass": ok} print(f"[3] McNemar vs scipy -> {'PASS' if ok else 'FAIL'}") return ok # ── 4. per-item schema contract ─────────────────────────────────────────────── def check_schema(): recs = [ {"item_id": "mmlu-0", "prompt_hash": "abc", "correct": True, "pred": "A", "gold": "A"}, {"item_id": "mmlu-1", "prompt_hash": "def", "correct": False, "pred": "B", "gold": "C"}, {"item_id": "mmlu-2", "prompt_hash": "ghi", "correct": True, "pred": "D", "gold": "D"}, ] fd, path = tempfile.mkstemp(suffix=".jsonl") with os.fdopen(fd, "w") as f: for r in recs: f.write(json.dumps(r) + "\n") read = mcnemar.read_items(path) os.unlink(path) required = {"item_id", "prompt_hash", "correct", "pred", "gold"} unique = len(read) == len(recs) complete = all(required <= set(v) for v in read.values()) typed = all(isinstance(v["correct"], bool) for v in read.values()) ok = unique and complete and typed results["schema_contract"] = {"unique_ids": unique, "fields_complete": complete, "correct_is_bool": typed, "pass": ok} print(f"[4] per-item schema: unique={unique} complete={complete} bool={typed} " f"-> {'PASS' if ok else 'FAIL'}") return ok def main(): checks = [check_eigh_cholesky, check_lm_head_guard, check_mcnemar, check_schema] passed = [c() for c in checks] root = os.path.dirname(os.path.abspath(__file__)) C.dump_json(results, os.path.join(root, "results", "stats", "selfcheck.json")) allok = all(passed) print(f"\n=== SELF-CHECK {'ALL PASS' if allok else 'FAIL'} " f"({sum(passed)}/{len(passed)}) ===") sys.exit(0 if allok else 1) if __name__ == "__main__": main()