1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
|
# zbp-scaling
Scaling study for **ZBP (zeroth-order backpropagation)** vs BP on decoder-only language models
(part 2 of the ZBP paper). Every nonlinear block is treated as a physical black box trained from
forward queries only (score-space attention core + per-token SwiGLU FFN); linear maps and the score
product are digital. See the main zobp repo for the method, theory (NSR ≈ c·d/n) and part-1/3 results.
## Layout
- `src/zbp_scaling/zbp/` — vendored ZBP package (probes, estimators, ZBPBlock autograd)
- `src/zbp_scaling/model.py` — OLMo2-style transformer (RMSNorm, SwiGLU, untied embeddings; learned
positions in v1) with the ZBP physical/digital partition
- `src/zbp_scaling/data.py` — FineWeb-Edu → GPT-2-BPE uint16 shards (WikiText-103 fallback for smoke)
- `src/zbp_scaling/trainer.py` — DDP (torchrun, 4 or 8 GPUs), bf16 autocast with fp32 measurement
accumulation, gradient accumulation to a fixed global batch, cosine + warmup, resume, JSONL logs
- `src/zbp_scaling/diagnostics.py` — per-block branch-gain rho_k profile and the NSR constant c(scale)
- `configs/` — model sizes (60m/124m/350m/1b) x training arms (bp / zbp_n16 / zbp_n4)
## Run
Data prep runs **on the training node** (H200), not on a dev machine; `data/` and `runs/` are gitignored.
```
python scripts/prepare_data.py --tokens 3e9 --out data/fineweb # on the H200 node
torchrun --nproc_per_node=8 scripts/train.py --model configs/model/m124.yaml --train configs/train/zbp_n16.yaml
```
Global batch is fixed in the train config; per-rank micro-batch and accumulation adapt to world size.
## Measured cost (48 GB Ampere-class, d=512 L=8 seq 1024, shared GPU)
BP 312 ms/step; ZBP n=16 **8.7x**, n=64 **29x** (FLOPs-bound: probe batching is already saturated at
probe_chunk 8; the score core uses its own chunk <= 4 since its memory goes as chunk*B*H*T^2).
Ladder arms: **bp / zbp_n16 / zbp_n64** per size (n=4 optional). Headroom if needed: torch.compile on the
query path; forward differences (n+1 instead of 2n queries) as a cheaper biased arm.
|