#!/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()