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