#!/usr/bin/env python3 """E3 Stage A: can protected learned bindings persist at all? Semantic tasks only (SEM_LATEST + SEM_2HOP), 4.39M, three seeds, two arms: --mode bind exact-address writes, unrelated slots bitwise untouched --mode unbound 192 slots, unstructured writes (E3-A3 capacity control) Arms are PAIRED: same seed gives identical initial parameters, identical lifetime and batch ordering, identical schedule and identical eval examples. Gates are frozen in preregistration/PREREGISTRATION_E3_BIND.md and evaluated by eval/stage_a_eval.py, not here. """ import argparse import json import sys import time from pathlib import Path import numpy as np import torch import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parent)) from pns.common import ckpt_root, logs_root, shards_root # noqa: E402 from pns.model.bind import BindConfig, PNSBind # noqa: E402 from pns.model.modules import enum_legal_mask # noqa: E402 from pns.train.loader import LifetimeBatcher, iter_eval_batches # noqa: E402 from pns.world.schema import ET, Fam # noqa: E402 SEM = (int(Fam.SEM_LATEST), int(Fam.SEM_2HOP)) TB = 32 def to_dev(b, dev): return {k: torch.from_numpy(np.ascontiguousarray( v.view(np.int64) if v.dtype == np.uint64 else v.astype(np.int64))).to(dev) for k, v in b.items()} @torch.no_grad() def evaluate(model, dev, split, n_lifetimes, intervention="none"): """Returns rows: (family, ok, delay, n_unrelated_writes_since_evidence).""" model.eval() rows = [] for b in iter_eval_batches(split, shards_root(), 24, n_lifetimes): g = to_dev(b, dev) B, L = g["etype"].shape state = model.initial_state(B, dev) if intervention == "payload_zero": state = torch.zeros_like(state) swap_at = L // 2 # count unrelated binding writes per event, for the interference split wcount = torch.zeros(B, device=dev) wcount_hist = torch.zeros(B, L, device=dev) with torch.autocast("cuda", dtype=torch.bfloat16): for t in range(L): if intervention == "payload_swap" and t == swap_at: state = state.roll(1, dims=0) state, out = model.step( state, 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=(intervention == "payload_zero")) wcount = wcount + (g["bind_write"][:, t] >= 0).float() wcount_hist[:, t] = wcount sel = torch.isin(g["family"][:, t], torch.tensor(SEM, device=dev)) if sel.any(): legal = enum_legal_mask(g["enum_legal"][sel, t]) pred = out["enum"][sel].masked_fill(~legal, -1e9).argmax(-1) ok = (pred == g["enum_gold"][sel, t]).float().cpu().numpy() idx = torch.nonzero(sel).flatten() for j, i in enumerate(idx.tolist()): d = int(g["delay"][i, t]) ev_t = max(0, t - d) n_unrel = float(wcount_hist[i, t] - wcount_hist[i, ev_t]) rows.append((int(g["family"][i, t]), float(ok[j]), d, n_unrel)) model.train() return rows def summarise(rows): a = np.array([r[1] for r in rows]) if rows else np.zeros(0) d = np.array([r[2] for r in rows]) if rows else np.zeros(0) u = np.array([r[3] for r in rows]) if rows else np.zeros(0) fam = np.array([r[0] for r in rows]) if rows else np.zeros(0) out = {"n": len(rows), "acc": float(a.mean()) if len(a) else float("nan")} for lo, hi, name in ((17, 10**9, "gt16"), (129, 10**9, "gt128"), (1, 16, "le16")): m = (d >= lo) & (d <= hi) if m.sum() > 20: out[f"acc_{name}"] = round(float(a[m].mean()), 4) out[f"n_{name}"] = int(m.sum()) for f, nm in ((int(Fam.SEM_LATEST), "latest"), (int(Fam.SEM_2HOP), "2hop")): m = fam == f if m.sum() > 20: out[f"acc_{nm}"] = round(float(a[m].mean()), 4) mm = m & (d > 16) if mm.sum() > 20: out[f"acc_{nm}_gt16"] = round(float(a[mm].mean()), 4) # interference split: accuracy vs unrelated writes since the evidence if len(u): q = np.quantile(u, [0.33, 0.66]) for i, (lo, hi) in enumerate([(-1, q[0]), (q[0], q[1]), (q[1], 1e9)]): m = (u > lo) & (u <= hi) if m.sum() > 20: out[f"acc_unrel_t{i}"] = round(float(a[m].mean()), 4) out[f"unrel_t{i}_mean"] = round(float(u[m].mean()), 1) return out def main(): ap = argparse.ArgumentParser() ap.add_argument("--mode", choices=["bind", "unbound"], required=True) ap.add_argument("--run", required=True) ap.add_argument("--seed", type=int, default=1) ap.add_argument("--updates", type=int, default=2600) ap.add_argument("--batch", type=int, default=128) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--eval-lifetimes", type=int, default=400) args = ap.parse_args() dev = "cuda" torch.manual_seed(args.seed) # identical init across arms model = PNSBind(BindConfig(mode=args.mode)).to(dev) opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) rd = ckpt_root() / args.run rd.mkdir(parents=True, exist_ok=True) log = open(logs_root() / f"{args.run}.jsonl", "a") print(f"{args.run}: mode={args.mode} seed={args.seed} " f"params={sum(p.numel() for p in model.parameters())/1e6:.2f}M", flush=True) batcher = LifetimeBatcher("e3_train", shards_root(), args.batch, seed=args.seed) # identical ordering across arms u, t0 = 0, time.time() for lt in batcher: if u >= args.updates: break g = to_dev(lt, dev) B, L = g["etype"].shape state = model.initial_state(B, dev) for c0 in range(0, L, TB): state = state.detach() opt.zero_grad(set_to_none=True) losses = [] with torch.autocast("cuda", dtype=torch.bfloat16): for t in range(c0, min(c0 + TB, L)): state, out = model.step( state, 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"]) sel = torch.isin(g["family"][:, t], torch.tensor(SEM, device=dev)) if sel.any(): legal = enum_legal_mask(g["enum_legal"][sel, t]) lg = out["enum"][sel].masked_fill(~legal, -1e9) losses.append(F.cross_entropy(lg, g["enum_gold"][sel, t])) if losses: loss = torch.stack(losses).mean() loss.backward() gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) if torch.isfinite(loss): opt.step() u += 1 if u % 200 == 0: s = {"u": u, "loss": round(float(loss), 4), "gn": round(float(gn), 2), "s": round(time.time() - t0)} log.write(json.dumps(s) + "\n"); log.flush() print(json.dumps(s), flush=True) if u >= args.updates: break torch.save({"model": model.state_dict(), "cfg": vars(model.cfg), "args": vars(args)}, rd / "final.pt") res = {"run": args.run, "mode": args.mode, "seed": args.seed} for iv in ("none", "payload_swap", "payload_zero"): res[iv] = summarise(evaluate(model, dev, "e3_dev", args.eval_lifetimes, iv)) print(iv, json.dumps(res[iv]), flush=True) res["delta_payload_swap"] = round(res["none"]["acc"] - res["payload_swap"]["acc"], 4) res["delta_payload_zero"] = round(res["none"]["acc"] - res["payload_zero"]["acc"], 4) from pns.common import atomic_write_json, eval_root atomic_write_json(eval_root() / f"STAGEA_{args.run}.json", res) print("delta_swap", res["delta_payload_swap"], "delta_zero", res["delta_payload_zero"], flush=True) if __name__ == "__main__": main()