pns-bind-25m / eval /confirmation_eval.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
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"
@torch.no_grad()
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()