traj-mc / code /selfcheck.py
ttishere's picture
Publish code/selfcheck.py
1d9a87f verified
Raw History Blame Contribute Delete
6.07 kB
"""
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()