Download FY2001/plugin_ppl_eval_reviewed.py from lfqian/annulus-year-plugins: direct link, hf CLI and curl.
- Browser
- Download file 3.97 kB
-
https://huggingface.co/lfqian/annulus-year-plugins/resolve/main/FY2001/plugin_ppl_eval_reviewed.py
- Command line
-
hf download hf://lfqian/annulus-year-plugins/FY2001/plugin_ppl_eval_reviewed.py
-
curl -L -o plugin_ppl_eval_reviewed.py https://huggingface.co/lfqian/annulus-year-plugins/resolve/main/FY2001/plugin_ppl_eval_reviewed.py
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) | |