论文笔记:Critical Batch Size Revisited
论文题目: Critical Batch Size Revisited: A Simple Empirical Approach to Large-Batch Language Model Training. 作者机构: Allen Institute for AI (William Merrill et al.) 一句话总结: 论文提出了一种通过'分支训练'直接测量临界 Batch Size (CBS) 的经验方法,发现现有的'梯度噪声尺度'方法在 LLM 训练中不可靠,并提出'Batch Size Warmup'策略,在不损失性能的前提下减少了 43% 的梯度步数。
1. 背景和动机
1.1 大模型训练的痛点与 Batch Size 权衡
- 背景: LLM 训练极其昂贵,提高吞吐量是核心诉求。
- 手段: 数据并行是主要手段,即增大 Batch Size (BS)。
- 权衡(Trade-off):
- BS 过小:训练慢,无法充分利用硬件并行能力。
- BS 过大:边际效应递减(Diminishing Returns)。虽然每步处理了更多数据,但模型收敛所需的 Token 总量变多了(样本效率下降)。
- 核心概念: Critical Batch Size (CBS, B*) —— 超过这个阈值,增加 BS 会导致计算效率下降。
1.2 既有方法及其局限性
- 主流理论: McCandlish et al. (2018) 提出的基于梯度噪声尺度 (Gradient Noise Scale) 的估算方法。
- 该理论认为:CBS = 梯度的方差与梯度范数之比。
- GPT-3 等著名工作都参考了这一理论。
- 本文的质疑: McCandlish 的方法依赖两个强假设:
- SGD 假设:假设优化器是 SGD(但 LLM 主要用 Adam)。
- 良态假设:假设 Hessian 矩阵是单位矩阵的倍数(实际并不成立)。
- 结论: 在 Adam 优化器和 LLM 场景下,噪声尺度(Noise Scale)可能不是 CBS 的有效代理。
2. 实验方法
2.1 本文提出的方法:分支训练
- 核心思想: 不依赖理论假设,直接用实验'测量'CBS。
- 操作步骤:
- 取一个训练中的检查点。
- 以当前 BS 为基准,开启多个'分支'训练任务。
- 每个分支使用不同的 BS 倍数,并相应调整学习率(Adam 用平方根缩放)。
- 训练一个小窗口步数 $\Delta$ (文中取 2B tokens)。
- 判定标准:如果大 BS 分支的 Loss 在 $\Delta$ 步后能恢复到与小 BS 分支相近(误差 $\epsilon$ 内),则认为该 BS 是'安全'的。
2.2 关键假设:局部恢复
- 假设: 如果在 $\Delta$ tokens 的短时间训练后,大 BS 的 Loss 能追平小 BS,那么在之后的训练中它也能保持住。
- 优势: 相比于 McCandlish 对优化器和 Loss 地形的强假设,这个'局部恢复'假设在工程上更弱、更易验证。
- 参数细节:
- 窗口 $\Delta$ = 2B tokens。
- 容忍度 $\epsilon$ = 0.01。
- Loss 经过了平滑处理。






