summaryrefslogtreecommitdiff
path: root/scripts/bench_probes.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/bench_probes.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/bench_probes.py')
-rw-r--r--scripts/bench_probes.py38
1 files changed, 38 insertions, 0 deletions
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")