diff options
| author | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:14:09 -0500 |
|---|---|---|
| committer | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:14:09 -0500 |
| commit | 6a544fabfc2af22e4d5823410dd2387b5af89ea9 (patch) | |
| tree | 0abd67bdda420deed27428b621fb59db8be07f41 /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.py | 97 |
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() |
