File size: 6,072 Bytes
1d9a87f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
"""
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()