"""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