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