From 6a544fabfc2af22e4d5823410dd2387b5af89ea9 Mon Sep 17 00:00:00 2001 From: yurenh Date: Mon, 31 Aug 2026 18:14:09 -0500 Subject: scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe --- scripts/bench_probes.py | 38 +++++++++++++++++++ scripts/train.py | 97 +++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 135 insertions(+) create mode 100644 scripts/bench_probes.py create mode 100644 scripts/train.py (limited to 'scripts') diff --git a/scripts/bench_probes.py b/scripts/bench_probes.py new file mode 100644 index 0000000..fc5faa0 --- /dev/null +++ b/scripts/bench_probes.py @@ -0,0 +1,38 @@ +"""Throughput vs probe_chunk for ZBP steps (informs the H200 config).""" +import os, sys, time, argparse +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) +import torch, torch.nn as nn +from zbp_scaling.model import ScalingLM +from zbp_scaling.zbp import ZBPConfig, zbp_blocks +from zbp_scaling.zbp.autograd import seed_probes + +p = argparse.ArgumentParser() +p.add_argument("--device", default="cuda:3") +p.add_argument("--d", type=int, default=512); p.add_argument("--layers", type=int, default=8) +p.add_argument("--vocab", type=int, default=8192); p.add_argument("--seq", type=int, default=1024) +p.add_argument("--bs", type=int, default=8) +a = p.parse_args() +dev = torch.device(a.device); seed_probes(0, dev) +x = torch.randint(0, a.vocab, (a.bs, a.seq), device=dev); y = torch.randint(0, a.vocab, (a.bs, a.seq), device=dev) + +def bench(mode, n, chunk, iters=4): + cfg = ZBPConfig(mode=mode, n_probes=n, eps=0.1, probe_chunk=chunk) + m = ScalingLM(a.vocab, a.d, a.layers, 8, a.seq, cfg=cfg).to(dev) + opt = torch.optim.AdamW(m.parameters(), lr=1e-4) + for _ in range(2): + loss = nn.functional.cross_entropy(m(x).flatten(0, 1), y.flatten()); opt.zero_grad(); loss.backward(); opt.step() + torch.cuda.synchronize(dev); t = time.time() + for _ in range(iters): + loss = nn.functional.cross_entropy(m(x).flatten(0, 1), y.flatten()); opt.zero_grad(); loss.backward(); opt.step() + torch.cuda.synchronize(dev) + dt = (time.time() - t) / iters + print(f"{mode:3s} n={n:3d} chunk={chunk:3d}: {dt*1000:7.0f} ms/step {a.bs*a.seq/dt/1000:7.1f} ktok/s peakmem {torch.cuda.max_memory_allocated(dev)/2**30:.1f} GB") + torch.cuda.reset_peak_memory_stats(dev) + return dt + +t_bp = bench("bp", 0, 0) +for n in (16, 64): + for chunk in (8, 16, 32, 64): + if chunk > 2 * n: continue + dt = bench("cd", n, chunk) + print(f" -> multiplier vs BP: {dt/t_bp:.1f}x") 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() -- cgit v1.2.3