批次大小缩放分析 | AI生成和翻译

首页 2026.07

让我并排比较三次运行。以下是增加批量大小时发生的情况:

  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 后发生了什么:

  1. 吞吐量保持不变 — 135K vs 129K tok/s。GPU 是计算受限而非内存受限,因此批量翻倍并不会使时间减半。步时间翻倍(4.1→7.8s),因为每一步处理的 token 数量翻倍。

  2. MFU 略有提升 — 65% vs 62%。更大的矩阵乘法在 H200 张量核心上效率稍高。

  3. 总步数减半 — 14,500 vs 29,000。相同总 token 数,步数更少。挂钟时间大致相同(~31h vs ~33h)。

  4. 梯度更干净 — 每步 1M token 而非 524K。梯度方差减半(Var ∝ 1/B)。自动缩放的学习率(×1.414)考虑了这一点——当梯度方向更确定时,可以迈出更大步。

  5. 学习率自动缩放 — nanochat 的 scaling law 代码检测到 2× 批量,并将学习率乘以 √2。这是正确的理论调整(线性缩放规则:LR ∝ batch size)。

为什么 64 和 128 会 OOM:

罪魁祸首是 lm_head — 一个 (32768 vocab × 1536 dim) 的线性层。在前向传播过程中,输出的 logits 是 (batch×seq, vocab)

lm_head 本身加上它的 fp32 优化器副本也消耗约 400 MB × 2 = 800 MB,但 logits 张量才是杀手。

结论: 将批量从 16 增加到 32 带来了更干净的梯度和略高的 MFU,但在 64 时遇到了 VRAM 瓶颈。H200 在 32K 词表和当前模型宽度下无法再提高。对于后面的收敛瓶颈阶段(大约第 10,000 步),我们将保持 batch=32——单张 H200 上没有进一步扩展的空间。


AI生成和翻译,但整理分享仍需精力。如果觉得有帮助,欢迎 捐助支持。 微信:@lzwjavaWeChat QR · X: @lzwjava · 打个招呼 👋

返回 捐助