File size: 3,966 Bytes
4886aa4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""FY2001 plugin PPL eval: backbone (OFF) vs backbone+plugin (ON) on held-out EDGAR 2001/2004.
Base = v11-33e-2001 with e32 masked (per CowBot repro config); plugin mounted @ L12,L23.
Same scorer as rung1 (doc_nll), prepend [Y2001]=151831. Outputs 4-cell PPL table."""
import os, sys, json, math, torch
S = "/gpfs/radev/scratch/xu_hua/lq62/annulus_v4"
PLUG = os.environ["PLUG_DIR"]          # dir with plugin_FY2001_r128_L12-23.pt + annulus_year_branch.py
sys.path.insert(0, S)
sys.path.insert(0, PLUG)
from rung1_2001_loader import build_rung1_model, doc_nll
import annulus_year_branch as AYB

CKPT = os.environ["CKPT"]; TOK = os.environ["TOK"]
YT = 151831  # [Y2001]
MAXTOK = 1024

print("[plugppl] building base v11-33e, then ZERO e32 weights (match plugin training backbone) ...", flush=True)
core, tok = build_rung1_model(CKPT, TOK)
# --- base construction: WEIGHT-ZERO e32 (exactly as year_branch_train.py:69, NOT router -inf) ---
zeroed = 0
for name, p in core.named_parameters():
    if ("local_experts.32.linear_fc" in name) or (name.endswith("weight32") and "mlp.experts.linear_fc" in name):
        p.data.zero_(); zeroed += 1
e32norm = 0.0
for name, p in core.named_parameters():
    if ("local_experts.32.linear_fc" in name) or (name.endswith("weight32") and "mlp.experts.linear_fc" in name):
        e32norm += float(p.data.norm().item())
print(f"[cert1] e32 weight-zero: zeroed {zeroed} tensors, e32_param_norm={e32norm:.6f} (expect 0)", flush=True)

# --- self-cert #2: OFF (branch disabled) must equal PURE base (no plugin) on a few docs ---
import json as _json
_probe = []
for _l in open(os.environ["DATA_2001"]):
    _t = tok(_json.loads(_l)["text"], add_special_tokens=False)["input_ids"][:1024]
    if len(_t) >= 8: _probe.append(_t)
    if len(_probe) >= 3: break
base_nll = [doc_nll(core, [151831], t)[0] for t in _probe]   # pure base, no plugin mounted yet

print("[plugppl] mounting plugin @ L12,L23 ...", flush=True)
branches, hooks = AYB.mount_plugin(core, os.path.join(PLUG, "plugin_FY2001_r128_L12-23.pt"),
                                   insert_layers=(12, 23))
for b in branches: b.enabled = False
off_nll = [doc_nll(core, [151831], t)[0] for t in _probe]     # OFF = branch disabled
maxdiff = max(abs(a - b) for a, b in zip(base_nll, off_nll))
print(f"[cert2] OFF==pure-base: max|NLL diff| over {len(_probe)} docs = {maxdiff:.2e} (expect ~0)", flush=True)
print(f"[plugppl] mounted {len(branches)} branches", flush=True)

def load(path, mx=200):
    out = []
    for line in open(path):
        d = json.loads(line)
        tids = tok(d["text"], add_special_tokens=False)["input_ids"][:MAXTOK]
        if len(tids) < 8:
            continue
        out.append(tids)
        if mx and len(out) >= mx:
            break
    return out

def score(docs):
    tn, tt = 0.0, 0
    for tids in docs:
        nll, n = doc_nll(core, [YT], tids)   # prepend [Y2001], score continuation
        tn += nll; tt += n
    m = tn / tt if tt else float("nan")
    return m, math.exp(m) if math.isfinite(m) else float("nan"), tt

DATA = {"2001": os.environ["DATA_2001"], "2004": os.environ["DATA_2004"]}
res = {}
for yr, path in DATA.items():
    docs = load(path, int(os.environ.get("N_DOCS", "120")))
    for state in ("OFF", "ON"):
        for b in branches:
            b.enabled = (state == "ON")
        nll, ppl, ntok = score(docs)
        res[f"{yr}-{state}"] = (nll, ppl, ntok, len(docs))
        print(f"[plugppl] {yr}-{state}: NLL={nll:.4f} PPL={ppl:.4f} ntok={ntok} ndoc={len(docs)}", flush=True)

print("PLUGIN_PPL_RESULT " + json.dumps({k: {"nll": v[0], "ppl": v[1], "ntok": v[2], "ndoc": v[3]}
                                         for k, v in res.items()}), flush=True)
# absorption + cutoff readout
for yr in ("2001", "2004"):
    off = res[f"{yr}-OFF"][0]; on = res[f"{yr}-ON"][0]
    print(f"[plugppl] {yr}: ΔNLL(ON-OFF)={on-off:+.4f} ({'降=帮' if on<off else '未降'})", flush=True)
print("PLUGIN_PPL_DONE", flush=True)