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