批次大小缩放分析 | AI生成和翻译
让我并排比较三次运行。以下是增加批量大小时发生的情况:
| batch=16 | batch=32 | batch=64 | batch=128 | |
|---|---|---|---|---|
| VRAM | 54 GB (38%) | 93 GB (65%) | OOM (133 GB) | OOM (139 GB) |
| Tok/sec | 129K | 135K | — | — |
| MFU | 62% | 65% | — | — |
| Step time | 4.1s | 7.8s | — | — |
| Grad accum | 16 | 16 | — | — |
| Batch/step | 524K | 1M | — | — |
| Total steps | 29,000 | 14,500 | — | — |
| ETA | ~33h | ~31h | — | — |
| LR | 1.0 | ×1.414 (auto-scaled) | — | — |
增加至 batch=32 后发生了什么:
-
吞吐量保持不变 — 135K vs 129K tok/s。GPU 是计算受限而非内存受限,因此批量翻倍并不会使时间减半。步时间翻倍(4.1→7.8s),因为每一步处理的 token 数量翻倍。
-
MFU 略有提升 — 65% vs 62%。更大的矩阵乘法在 H200 张量核心上效率稍高。
-
总步数减半 — 14,500 vs 29,000。相同总 token 数,步数更少。挂钟时间大致相同(~31h vs ~33h)。
-
梯度更干净 — 每步 1M token 而非 524K。梯度方差减半(Var ∝ 1/B)。自动缩放的学习率(×1.414)考虑了这一点——当梯度方向更确定时,可以迈出更大步。
-
学习率自动缩放 — nanochat 的 scaling law 代码检测到 2× 批量,并将学习率乘以 √2。这是正确的理论调整(线性缩放规则:LR ∝ batch size)。
为什么 64 和 128 会 OOM:
罪魁祸首是 lm_head — 一个 (32768 vocab × 1536 dim) 的线性层。在前向传播过程中,输出的 logits 是 (batch×seq, vocab):
- batch=32: (65536, 32768) 在 fp32 中 = 8 GB → 总计 93 GB ✅
- batch=64: (131072, 32768) 在 fp32 中 = 16 GB → 总计 133 GB OOM ❌
- batch=128: (262144, 32768) 在 fp32 中 = 32 GB → 总计 139 GB OOM ❌
lm_head 本身加上它的 fp32 优化器副本也消耗约 400 MB × 2 = 800 MB,但 logits 张量才是杀手。
结论: 将批量从 16 增加到 32 带来了更干净的梯度和略高的 MFU,但在 64 时遇到了 VRAM 瓶颈。H200 在 32K 词表和当前模型宽度下无法再提高。对于后面的收敛瓶颈阶段(大约第 10,000 步),我们将保持 batch=32——单张 H200 上没有进一步扩展的空间。
