"""torchrun-able BP/ZBP trainer. Usage: torchrun --nproc_per_node=8 scripts/train.py --model configs/model/m124.yaml --train configs/train/zbp_n16.yaml \ --data data/fineweb --out runs/m124_zbp16 [--set key=value ...] Single-GPU: python scripts/train.py ... [--device cuda:0]""" import os, sys, json, math, time, argparse sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) import yaml import torch import torch.nn as nn import torch.distributed as dist from zbp_scaling.model import ScalingLM from zbp_scaling.data import Shards from zbp_scaling.zbp import ZBPConfig, zbp_blocks from zbp_scaling.zbp.autograd import seed_probes p = argparse.ArgumentParser() p.add_argument("--model", required=True) p.add_argument("--train", required=True) p.add_argument("--data", required=True) p.add_argument("--out", default="runs/dev") p.add_argument("--device", default=None) p.add_argument("--steps", type=int, default=None, help="override: stop after this many steps (smoke)") p.add_argument("--micro_bs", type=int, default=8) p.add_argument("--set", nargs="*", default=[], help="key=value overrides for the yaml configs") a = p.parse_args() M = yaml.safe_load(open(a.model)); Tr = yaml.safe_load(open(a.train)) for kv in a.set: k, v = kv.split("=", 1) tgt = M if k in M else Tr tgt[k] = yaml.safe_load(v) ddp = "WORLD_SIZE" in os.environ and int(os.environ["WORLD_SIZE"]) > 1 rank, world = (int(os.environ.get("RANK", 0)), int(os.environ.get("WORLD_SIZE", 1))) if ddp: dist.init_process_group("nccl") dev = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}") torch.cuda.set_device(dev) else: dev = torch.device(a.device or ("cuda:0" if torch.cuda.is_available() else "cpu")) torch.manual_seed(Tr.get("seed", 0) + rank) seed_probes(Tr.get("seed", 0) * 1000 + rank, dev) cfg = ZBPConfig(mode=Tr["mode"], n_probes=Tr.get("n_probes", 4), probe=Tr.get("probe", "rademacher"), eps=Tr.get("eps", 0.1), probe_chunk=Tr.get("probe_chunk", 8)) model = ScalingLM(M["vocab"], M["d"], M["layers"], M["heads"], M["seq_len"], M.get("ffn_mult", "8/3"), cfg).to(dev) raw = model if ddp: model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[dev.index]) n_params = sum(q.numel() for q in raw.parameters()) seq = M["seq_len"] accum = max(1, int(Tr["global_batch_tokens"]) // (seq * a.micro_bs * world)) total_steps = a.steps or int(float(Tr["tokens"]) // int(Tr["global_batch_tokens"])) opt = torch.optim.AdamW(raw.parameters(), lr=float(Tr["lr"]), weight_decay=Tr.get("wd", 0.1), betas=tuple(Tr.get("betas", [0.9, 0.95]))) sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / Tr.get("warmup", 2000)) * 0.5 * (1 + math.cos(math.pi * min(s, total_steps) / total_steps))) data = Shards(a.data, seq, dev) gen = torch.Generator(); gen.manual_seed(Tr.get("seed", 0) * 77 + rank) os.makedirs(a.out, exist_ok=True) logf = open(os.path.join(a.out, f"log_rank{rank}.jsonl"), "a") def log(**kw): if rank == 0: kw["time"] = time.time() logf.write(json.dumps(kw) + "\n"); logf.flush(); print(json.dumps(kw), flush=True) @torch.no_grad() def evaluate(nb=20): model.eval(); tot = 0.0 for _ in range(nb): x, y = data.batch("val", a.micro_bs, gen) tot += nn.functional.cross_entropy(model(x).flatten(0, 1), y.flatten()).item() model.train(); return tot / nb if rank == 0: log(kind="meta", params=n_params, accum=accum, world=world, total_steps=total_steps, cfg=cfg.asdict(), model=M, train={k: str(v) for k, v in Tr.items()}) t0 = time.time() for step in range(total_steps + 1): if step % Tr.get("eval_every", 500) == 0 or step == total_steps: log(kind="eval", step=step, val_loss=evaluate(), elapsed=time.time() - t0, tok_per_s=step * int(Tr["global_batch_tokens"]) / max(time.time() - t0, 1e-9)) if rank == 0 and step % (10 * Tr.get("eval_every", 500)) == 0: torch.save({"model": raw.state_dict(), "opt": opt.state_dict(), "step": step}, os.path.join(a.out, "ckpt.pt")) if step == total_steps: break opt.zero_grad(set_to_none=True) for micro in range(accum): x, y = data.batch("train", a.micro_bs, gen) loss = nn.functional.cross_entropy(model(x).flatten(0, 1), y.flatten()) / accum loss.backward() if Tr.get("grad_clip", 0): torch.nn.utils.clip_grad_norm_(raw.parameters(), Tr["grad_clip"]) opt.step(); sched.step() if step % 50 == 0: log(kind="train", step=step, loss=loss.item() * accum, lr=sched.get_last_lr()[0]) if ddp: dist.destroy_process_group()