pns-bind-25m / src /stage_b.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
5.42 kB
#!/usr/bin/env python3
"""E3 Stage B: the FULL multitask mixture that defeated all seven prior models.
Same corpus, same objective mixture, same head system, same budget and the
same loss weights as the failed runs. The ONLY difference is the state
mechanism (addressed binding writes vs unstructured). Per the frozen
preregistration: no enum weighting, no continuation, no retention tricks, no
post-hoc tuning.
Parameter-matched: enc_layers=8 -> 24.27M vs PNSR's 24.61M (-1.4%, inside the
+/-3% policy).
"""
import argparse
import json
import sys
import time
from pathlib import Path
import numpy as np
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent))
from pns.common import atomic_write_json, ckpt_root, eval_root, logs_root, shards_root # noqa
from pns.model import modules # noqa: E402
from pns.model.bind import BindConfig, PNSBind # noqa: E402
from pns.train.loader import LifetimeBatcher # noqa: E402
TBPTT = 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 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=6000)
ap.add_argument("--batch", type=int, default=96)
ap.add_argument("--lr", type=float, default=4e-4)
args = ap.parse_args()
dev = "cuda"
torch.manual_seed(args.seed)
cfg = BindConfig(d=384, heads=8, enc_layers=8, mode=args.mode, full_heads=True)
model = PNSBind(cfg).to(dev)
opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95),
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} params="
f"{sum(p.numel() for p in model.parameters())/1e6:.2f}M", flush=True)
batcher = LifetimeBatcher("e3_train", shards_root(), args.batch, seed=args.seed)
d = cfg.d
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, TBPTT):
if u >= args.updates:
break
c1 = min(c0 + TBPTT, L)
state = state.detach()
for pg in opt.param_groups:
pg["lr"] = args.lr * (min(1.0, (u + 1) / 300) if u < 300
else (0.1 + 0.9 * max(0.0, 1 - (u - 300)
/ max(1, args.updates - 300))))
opt.zero_grad(set_to_none=True)
outs = {k: [] for k in ("mode", "enum", "ptr", "op", "args")}
with torch.autocast("cuda", dtype=torch.bfloat16):
h_static = model.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(c0, c1):
live = g["live"][:, t]
mask = live >= 0
rows = live.clamp(min=0)
bank = torch.gather(h_static, 1, rows.unsqueeze(-1).expand(-1, -1, d))
bank = model.recenc.finalize(
bank, (t - torch.gather(g["rec_birth"], 1, rows)).clamp(min=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"],
rec_bank=bank * mask.unsqueeze(-1), live_mask=mask)
for k in outs:
outs[k].append(out[k])
flat = {k: torch.cat([o.unsqueeze(1) for o in v], 1).flatten(0, 1)
for k, v in outs.items()}
sup = {k: g[k][:, c0:c1].flatten(0, 1) for k in
("mode_gold", "enum_gold", "ptr_gold_slot", "op_gold")}
sup["enum_legal"] = g["enum_legal"][:, c0:c1].flatten(0, 1)
sup["op_arg_slots"] = g["op_arg_slots"][:, c0:c1].flatten(0, 1)
# IDENTICAL loss weights to the failed runs (w_enum default 2.0)
L_parts = modules.losses(flat, sup)
loss = sum(L_parts.values())
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)
torch.save({"model": model.state_dict(), "cfg": vars(cfg), "args": vars(args)},
rd / "final.pt")
atomic_write_json(eval_root() / f"STAGEB_{args.run}_done.json",
{"run": args.run, "updates": u})
print(f"done {u} updates", flush=True)
if __name__ == "__main__":
main()