#!/usr/bin/env python3 """Train one model of the preregistered zoo on the frozen PNS corpus. --model pnsr | pnsr_k1 | tx768 | txe Single GPU (selected by UUID via CUDA_VISIBLE_DEVICES in the launch script). PNSR: streaming TBPTT over whole lifetimes, state detached NEVER reset. TX: supervised-position window sampling, exposure-matched to PNSR. """ import argparse import json import math import os 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 ckpt_root, logs_root, shards_root # noqa: E402 from pns.model import modules # noqa: E402 from pns.model.pnsr import PNSR, PNSRConfig # noqa: E402 from pns.model.rmt import RMT, RMTConfig # noqa: E402 from pns.model.tx import TX, TXConfig # noqa: E402 from pns.train.loader import LifetimeBatcher, WindowSampler # noqa: E402 TBPTT = 32 def atomic_save(obj, path: Path): tmp = path.with_suffix(".tmp") torch.save(obj, tmp) with open(tmp, "rb") as f: os.fsync(f.fileno()) os.replace(tmp, path) def gather_bank(model, h_static, live, rec_birth, ev_idx, d): """live: [B,100] int (record row or -1) -> (bank [B,100,d], mask [B,100]).""" mask = live >= 0 rows = live.clamp(min=0).long() bank = torch.gather(h_static, 1, rows.unsqueeze(-1).expand(-1, -1, d)) births = torch.gather(rec_birth, 1, rows) age = (ev_idx - births).clamp(min=0) bank = model.recenc.finalize(bank, age) return bank * mask.unsqueeze(-1), mask def build_sup(flat): keys = ("mode_gold", "enum_gold", "enum_legal", "ptr_gold_slot", "op_gold", "op_arg_slots") return {k: flat[k] for k in keys} def lr_at(u, updates, base, warmup=300, hold_frac=0.7): """v2 recipe: hold base lr until hold_frac of training (algorithmic subtasks like value comparison have a long plateau before their transition - measured in the IMMEDIATE_CMP probe), then cosine to 0.1x.""" if u < warmup: return base * (u + 1) / warmup hold_end = hold_frac * updates if u < hold_end: return base p = (u - hold_end) / max(1.0, updates - hold_end) return base * (0.1 + 0.45 * (1 + math.cos(math.pi * p))) def train_pnsr(args, K): """Streaming TBPTT trainer. Works for any model exposing initial_state()/step() - PNSR and the RMT memory-token baseline.""" dev = "cuda" torch.manual_seed(args.seed) if args.model == "rmt": cfg = RMTConfig() model = RMT(cfg).to(dev) else: cfg = PNSRConfig(K=K, update_rule=args.update_rule, state_tau=args.state_tau) model = PNSR(cfg).to(dev) torch.nn.init.constant_(model.gate[-1].bias, args.gate_bias) n_par = sum(p.numel() for p in model.parameters()) opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.01) run_dir = ckpt_root() / args.run run_dir.mkdir(parents=True, exist_ok=True) log = open(logs_root() / f"{args.run}.jsonl", "a") start_u = 0 if args.resume_from: src = torch.load(ckpt_root() / args.resume_from / "final.pt", map_location=dev, weights_only=False) model.load_state_dict(src["model"]) print(f"initialized weights from {args.resume_from}/final.pt " f"(u={src['update']})", flush=True) if (run_dir / "latest.pt").exists(): ck = torch.load(run_dir / "latest.pt", map_location=dev, weights_only=False) model.load_state_dict(ck["model"]) opt.load_state_dict(ck["opt"]) start_u = ck["update"] print(f"resumed at update {start_u}", flush=True) batcher = LifetimeBatcher(args.shards_split, shards_root(), args.batch, seed=args.seed) d = cfg.d u = start_u t0, tok_count = time.time(), 0 torch.backends.cuda.matmul.allow_tf32 = True for lt in batcher: if u >= args.updates: break g = {k: torch.from_numpy(np.ascontiguousarray(v.astype(np.int64)) if v.dtype != np.uint16 else np.ascontiguousarray(v.astype(np.int64))).to(dev) for k, v in lt.items()} 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"] = lr_at(u, args.updates, args.lr, hold_frac=0.7 if args.hold_frac is None else args.hold_frac) 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): bank, mask = gather_bank(model, h_static, g["live"][:, t], g["rec_birth"], t, d) state, out = model.step(state, g["tok"][:, t], g["etype"][:, t], g["dt"][:, t], bank, mask) for k in outs: outs[k].append(out[k]) flat_out = {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) L_parts = modules.losses(flat_out, sup, w_enum=args.w_enum, w_other=args.w_other) loss = sum(L_parts.values()) loss.backward() gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) if torch.isfinite(loss): # clip_grad_norm_ already bounds the step opt.step() u += 1 tok_count += int(B * (c1 - c0)) if u % 20 == 0: with torch.no_grad(): mg = sup["mode_gold"] stats = {"u": u, "loss": round(float(loss), 4), "gn": round(float(gn), 2), "lr": round(opt.param_groups[0]["lr"], 6), "ev_s": round(tok_count / (time.time() - t0)), "mem_gb": round(torch.cuda.max_memory_allocated() / 2**30, 1), "state_norm": round(float(state.norm(dim=-1).mean()), 2)} from pns.model.modules import enum_legal_mask for name, mid in (("ptr", 1), ("enum", 2), ("op", 3)): m = mg == mid if m.any(): logits = flat_out[name][m] if name == "enum": legal = enum_legal_mask(sup["enum_legal"][m]) logits = logits.masked_fill(~legal, float("-inf")) pred = logits.argmax(-1) gold = sup[{"ptr": "ptr_gold_slot", "enum": "enum_gold", "op": "op_gold"}[name]][m] stats[f"acc_{name}"] = round(float((pred == gold).float().mean()), 3) for k, v in L_parts.items(): stats[f"L_{k}"] = round(float(v), 4) log.write(json.dumps(stats) + "\n") log.flush() if u % 200 == 0: print(json.dumps(stats), flush=True) if u % 500 == 0 or u == args.updates: atomic_save({"model": model.state_dict(), "opt": opt.state_dict(), "update": u, "cfg": vars(cfg), "params": n_par, "args": vars(args)}, run_dir / "latest.pt") atomic_save({"model": model.state_dict(), "update": u, "cfg": vars(cfg), "params": n_par, "args": vars(args)}, run_dir / "final.pt") print(f"done: {u} updates, params {n_par / 1e6:.2f}M", flush=True) def train_tx(args, window): dev = "cuda" torch.manual_seed(args.seed) cfg = TXConfig(window=window) model = TX(cfg).to(dev) n_par = sum(p.numel() for p in model.parameters()) opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.01) run_dir = ckpt_root() / args.run run_dir.mkdir(parents=True, exist_ok=True) log = open(logs_root() / f"{args.run}.jsonl", "a") start_u = 0 if args.resume_from: src = torch.load(ckpt_root() / args.resume_from / "final.pt", map_location=dev, weights_only=False) model.load_state_dict(src["model"]) print(f"initialized weights from {args.resume_from}/final.pt " f"(u={src['update']})", flush=True) if (run_dir / "latest.pt").exists(): ck = torch.load(run_dir / "latest.pt", map_location=dev, weights_only=False) model.load_state_dict(ck["model"]) opt.load_state_dict(ck["opt"]) start_u = ck["update"] print(f"resumed at update {start_u}", flush=True) sampler = WindowSampler(args.shards_split, shards_root(), args.batch, window, seed=args.seed) d = cfg.d u, t0, npos = start_u, time.time(), 0 recent, best = [], float("inf") for b in sampler: if u >= args.updates: break g = {k: torch.from_numpy(np.ascontiguousarray(v.astype(np.int64))).to(dev) for k, v in b.items() if k != "rec_val_hash"} for pg in opt.param_groups: pg["lr"] = lr_at(u, args.updates, args.lr, hold_frac=0.5 if args.hold_frac is None else args.hold_frac) opt.zero_grad(set_to_none=True) 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"]) bank, mask = gather_bank(model, h_static, g["live"], g["rec_birth"], g["ev_idx"].unsqueeze(1), d) out = model(g["tok"], bank, mask) L_parts = modules.losses(out, build_sup(g), w_enum=args.w_enum, w_other=args.w_other) loss = sum(L_parts.values()) loss.backward() gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) if torch.isfinite(loss): # clip_grad_norm_ already bounds the step opt.step() u += 1 npos += g["tok"].shape[0] recent.append(float(loss)) if len(recent) > 200: recent.pop(0) if u % 250 == 0 and len(recent) >= 200: import statistics med = statistics.median(recent) if med < best: best = med atomic_save({"model": model.state_dict(), "update": u, "cfg": vars(cfg), "params": n_par, "args": vars(args), "median_loss": med}, run_dir / "best.pt") if u % 50 == 0: with torch.no_grad(): stats = {"u": u, "loss": round(float(loss), 4), "gn": round(float(gn), 2), "pos_s": round(npos / (time.time() - t0)), "mem_gb": round(torch.cuda.max_memory_allocated() / 2**30, 1)} from pns.model.modules import enum_legal_mask mg = g["mode_gold"] for name, mid in (("ptr", 1), ("enum", 2)): m = mg == mid if m.any(): logits = out[name][m] if name == "enum": legal = enum_legal_mask(g["enum_legal"][m]) logits = logits.masked_fill(~legal, float("-inf")) pred = logits.argmax(-1) gold = g[{"ptr": "ptr_gold_slot", "enum": "enum_gold"}[name]][m] stats[f"acc_{name}"] = round(float((pred == gold).float().mean()), 3) log.write(json.dumps(stats) + "\n") log.flush() if u % 500 == 0: print(json.dumps(stats), flush=True) if u % 1000 == 0 or u == args.updates: atomic_save({"model": model.state_dict(), "opt": opt.state_dict(), "update": u, "cfg": vars(cfg), "params": n_par, "args": vars(args)}, run_dir / "latest.pt") atomic_save({"model": model.state_dict(), "update": u, "cfg": vars(cfg), "params": n_par, "args": vars(args)}, run_dir / "final.pt") print(f"done: {u} updates, params {n_par / 1e6:.2f}M", flush=True) def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", required=True, choices=["pnsr", "pnsr_k1", "tx768", "txe", "rmt"]) ap.add_argument("--run", required=True) ap.add_argument("--updates", type=int, default=None) ap.add_argument("--batch", type=int, default=None) ap.add_argument("--lr", type=float, default=None) ap.add_argument("--seed", type=int, default=1) ap.add_argument("--resume-from", default=None, help="run name whose final.pt initializes the weights (fresh opt)") ap.add_argument("--w-enum", type=float, default=2.0) ap.add_argument("--w-other", type=float, default=1.0) ap.add_argument("--update-rule", default="additive_clamp", choices=["additive_clamp", "convex"]) ap.add_argument("--gate-bias", type=float, default=0.0) ap.add_argument("--state-tau", type=float, default=16.0) ap.add_argument("--shards-split", default="train") ap.add_argument("--hold-frac", type=float, default=None, help="lr hold fraction; default 0.7 pnsr / 0.5 tx; 0 = decay from start") args = ap.parse_args() if args.lr is None: # TX class diverged at held 4e-4 (v2 s1 restart evidence); 3e-4 matches # the probe-validated constant-lr regime. PNSR stable at 4e-4. args.lr = 4e-4 if args.model.startswith("pnsr") else 3e-4 if args.model in ("pnsr", "pnsr_k1", "rmt"): args.updates = args.updates or 6000 args.batch = args.batch or 192 train_pnsr(args, K=1 if args.model == "pnsr_k1" else 4) else: args.updates = args.updates or 30000 args.batch = args.batch or 384 train_tx(args, window=768 if args.model == "tx768" else 64) if __name__ == "__main__": main()