summaryrefslogtreecommitdiff
path: root/scripts/train.py
diff options
context:
space:
mode:
authoryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
committeryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
commit6a544fabfc2af22e4d5823410dd2387b5af89ea9 (patch)
tree0abd67bdda420deed27428b621fb59db8be07f41 /scripts/train.py
scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe
Diffstat (limited to 'scripts/train.py')
-rw-r--r--scripts/train.py97
1 files changed, 97 insertions, 0 deletions
diff --git a/scripts/train.py b/scripts/train.py
new file mode 100644
index 0000000..af43918
--- /dev/null
+++ b/scripts/train.py
@@ -0,0 +1,97 @@
+"""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()