annulus-year-plugins / FY2001 /plugin_ppl_eval_reviewed.py
lfqian's picture
Upload FY2001/plugin_ppl_eval_reviewed.py with huggingface_hub
4886aa4 verified
Raw History Blame Contribute Delete
3.97 kB
"""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)