Download src/stage_a.py from nur-dev/pns-bind-25m: direct link, hf CLI and curl.
- Browser
- Download file 8.45 kB
-
https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/src/stage_a.py
- Command line
-
hf download hf://nur-dev/pns-bind-25m/src/stage_a.py
-
curl -L -o stage_a.py https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/src/stage_a.py
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()} | |
| 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() | |