Download eval/confirmation_eval.py from nur-dev/pns-bind-25m: direct link, hf CLI and curl.
- Browser
- Download file 8.3 kB
-
https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/confirmation_eval.py
- Command line
-
hf download hf://nur-dev/pns-bind-25m/eval/confirmation_eval.py
-
curl -L -o confirmation_eval.py https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/eval/confirmation_eval.py
8.3 kB
| #!/usr/bin/env python3 | |
| """Experiment 2 (PNS-Bind) evaluation and preregistered adjudication. | |
| This is the study's sealed-confirmation evaluator with the single-use sentinel | |
| removed: the confirmation split was read exactly once during the study, that | |
| use is recorded in `results/sealed_use_record.json`, and the split is public | |
| now, so the refusal machinery would only block reproduction. | |
| Frozen conjunction (preregistration/PREREGISTRATION_E3_BIND.md), required for EVERY seed: | |
| SEM_LATEST >= TXE + 0.05 AND SEM_2HOP >= TXE + 0.05 | |
| pooled mean >= TXE + 0.08 | |
| delta payload-swap >= 0.10 | |
| delta reset64 >= 0.10 | |
| Leakage control: payload-zero must collapse accuracy, i.e. the exact address | |
| alone must not carry the answer. | |
| python3 eval/confirmation_eval.py --split e3_conf | |
| """ | |
| import argparse | |
| import json | |
| import subprocess | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| ROOT = Path(__file__).resolve().parent.parent | |
| sys.path.insert(0, str(ROOT / "src")) | |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) | |
| from pns.checkpoint import load_model # noqa: E402 | |
| from pns.common import atomic_write_json, eval_root, shards_root # noqa: E402 | |
| from pns.model.modules import enum_legal_mask # noqa: E402 | |
| from pns.train.loader import iter_eval_batches # noqa: E402 | |
| from pns.world.schema import Fam # noqa: E402 | |
| from stage_a import to_dev # noqa: E402 | |
| SEM = (int(Fam.SEM_LATEST), int(Fam.SEM_2HOP)) | |
| INTERVENTIONS = ("none", "payload_swap", "reset64", "payload_zero") | |
| def cache_path(name, split, interv, limit): | |
| suf = "" if limit is None else f"_n{limit}" | |
| return eval_root() / f"E2_{name}_{split}_{interv}{suf}.json" | |
| def run(name, dev, split, interv="none", limit=None, batch=24, use_cache=True): | |
| """Evaluate one (checkpoint, intervention) pair. Cached on disk so the | |
| battery is resumable and can be sharded across devices.""" | |
| cp = cache_path(name, split, interv, limit) | |
| if use_cache and cp.exists(): | |
| return json.loads(cp.read_text()) | |
| m, _, _ = load_model(name, dev) | |
| d = m.cfg.d | |
| rows = [] | |
| for b in iter_eval_batches(split, shards_root(), batch, limit): | |
| g = to_dev(b, dev) | |
| B, L = g["etype"].shape | |
| st = m.initial_state(B, dev) | |
| if interv == "payload_zero": | |
| st = torch.zeros_like(st) | |
| swap_at = L // 2 | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| h = m.recenc.static(g["rec_val_toks"], g["rec_key_toks"], g["rec_store"], | |
| g["rec_kind"], g["rec_key"], g["rec_ent"]) | |
| for t in range(L): | |
| if interv == "payload_swap" and t == swap_at: | |
| st = st.roll(1, dims=0) | |
| if interv == "reset64" and t > 0 and t % 64 == 0: | |
| st = m.initial_state(B, dev) | |
| live = g["live"][:, t] | |
| mask = live >= 0 | |
| rr = live.clamp(min=0) | |
| bank = torch.gather(h, 1, rr.unsqueeze(-1).expand(-1, -1, d)) | |
| bank = m.recenc.finalize( | |
| bank, (t - torch.gather(g["rec_birth"], 1, rr)).clamp(min=0)) | |
| st, o = m.step(st, g["tok"][:, t], g["etype"][:, t], g["dt"][:, t], | |
| g["bind_write"][:, t], g["bind_read"][:, t], | |
| g["bind_slot_ent"], g["bind_slot_attr"], | |
| freeze_writes=(interv == "payload_zero"), | |
| rec_bank=bank * mask.unsqueeze(-1), live_mask=mask) | |
| sel = torch.isin(g["family"][:, t], torch.tensor(SEM, device=dev)) | |
| if sel.any(): | |
| lg = o["enum"][sel].masked_fill( | |
| ~enum_legal_mask(g["enum_legal"][sel, t]), -1e9) | |
| ok = (lg.argmax(-1) == g["enum_gold"][sel, t]).cpu().numpy() | |
| for j, i in enumerate(torch.nonzero(sel).flatten().tolist()): | |
| rows.append((int(g["family"][i, t]), float(ok[j]), | |
| int(g["delay"][i, t]))) | |
| a = np.array(rows) | |
| out = {"n": len(a), "acc": round(float(a[:, 1].mean()), 4)} | |
| for f, nm in ((int(Fam.SEM_LATEST), "latest"), (int(Fam.SEM_2HOP), "2hop")): | |
| mm = a[:, 0] == f | |
| if mm.sum() > 20: | |
| out[f"acc_{nm}"] = round(float(a[mm, 1].mean()), 4) | |
| out[f"n_{nm}"] = int(mm.sum()) | |
| for lo, nm in ((16, "gt16"), (128, "gt128")): | |
| mm = a[:, 2] > lo | |
| if mm.sum() > 20: | |
| out[f"acc_{nm}"] = round(float(a[mm, 1].mean()), 4) | |
| out[f"n_{nm}"] = int(mm.sum()) | |
| atomic_write_json(cp, out) | |
| return out | |
| def txe_comparator(split, batch): | |
| """Memory-free control on the SAME split, via the shared evaluator.""" | |
| tag = f"TXE_ON_E3_{split}_none" | |
| rowfile = eval_root() / f"rows_{tag}.npz" | |
| if not rowfile.exists(): | |
| subprocess.run([sys.executable, str(Path(__file__).with_name("evaluate.py")), | |
| "--run", "TXE_ON_E3", "--split", split, | |
| "--batch", str(batch)], check=True) | |
| z = np.load(rowfile) | |
| fam, ok = z["fam"], z["ok"] | |
| return {"latest": round(float(ok[fam == int(Fam.SEM_LATEST)].mean()), 4), | |
| "2hop": round(float(ok[fam == int(Fam.SEM_2HOP)].mean()), 4), | |
| "pooled": round(float(ok[np.isin(fam, SEM)].mean()), 4)} | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--split", default="e3_conf", | |
| help="e3_conf (the preregistered confirmation split) or e3_dev") | |
| ap.add_argument("--stage", default="B", choices=["B"], | |
| help="Stage B checkpoints (E3B_*)") | |
| ap.add_argument("--limit", type=int, default=None) | |
| ap.add_argument("--batch", type=int, default=24) | |
| ap.add_argument("--txe-batch", type=int, default=256) | |
| ap.add_argument("--out", default=None) | |
| ap.add_argument("--only", default=None, | |
| help="evaluate just this run (shard the battery), then exit") | |
| ap.add_argument("--interv", default=None, | |
| help="with --only: just this intervention") | |
| args = ap.parse_args() | |
| dev = "cuda" | |
| if args.only: | |
| for iv in ([args.interv] if args.interv else INTERVENTIONS): | |
| r = run(args.only, dev, args.split, iv, args.limit, args.batch) | |
| print(args.only, iv, json.dumps(r), flush=True) | |
| return | |
| rep = {"split": args.split, "TXE": txe_comparator(args.split, args.txe_batch)} | |
| T = rep["TXE"] | |
| for mode in ("bind", "unbound"): | |
| for s in (1, 2, 3): | |
| nm = f"E3B_{mode}_s{s}" | |
| r = {iv: run(nm, dev, args.split, iv, args.limit, args.batch) | |
| for iv in INTERVENTIONS} | |
| r["delta_payload_swap"] = round(r["none"]["acc"] - r["payload_swap"]["acc"], 4) | |
| r["delta_reset64"] = round(r["none"]["acc"] - r["reset64"]["acc"], 4) | |
| r["delta_payload_zero"] = round(r["none"]["acc"] - r["payload_zero"]["acc"], 4) | |
| rep[nm] = r | |
| print(nm, "acc", r["none"]["acc"], "latest", r["none"].get("acc_latest"), | |
| "2hop", r["none"].get("acc_2hop"), "dswap", r["delta_payload_swap"], | |
| "dreset", r["delta_reset64"], flush=True) | |
| v, seeds_ok = {}, [] | |
| for s in (1, 2, 3): | |
| r = rep[f"E3B_bind_s{s}"] | |
| c = {"latest": r["none"].get("acc_latest", 0) >= T["latest"] + 0.05, | |
| "2hop": r["none"].get("acc_2hop", 0) >= T["2hop"] + 0.05, | |
| "swap": r["delta_payload_swap"] >= 0.10, | |
| "reset": r["delta_reset64"] >= 0.10} | |
| c["all"] = all(c.values()) | |
| seeds_ok.append(c["all"]) | |
| v[f"seed{s}"] = c | |
| pooled = float(np.mean([rep[f"E3B_bind_s{s}"]["none"]["acc"] for s in (1, 2, 3)])) | |
| v["pooled_gain"] = round(pooled - T["pooled"], 4) | |
| v["pooled_ok"] = bool(v["pooled_gain"] >= 0.08) | |
| v["CONJUNCTION_MET"] = bool(all(seeds_ok) and v["pooled_ok"]) | |
| rep["verdict"] = v | |
| out = Path(args.out) if args.out else eval_root() / f"CONFIRM_{args.split}.json" | |
| atomic_write_json(out, rep) | |
| print("\nCONJUNCTION MET:", v["CONJUNCTION_MET"], "| pooled gain", v["pooled_gain"]) | |
| for s in (1, 2, 3): | |
| print(f" seed{s}:", json.dumps(v[f"seed{s}"])) | |
| print("->", out) | |
| if __name__ == "__main__": | |
| main() | |