pns-bind-25m / src /stage_a.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
8.45 kB
#!/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()