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