← 返回题目列表

大模型预训练 Batch Size 如何影响收敛?

中等 第 14 / 25 题 更新于 2026/09/18
大模型预训练Batch Size梯度累积学习率

简化版

更大 Batch 会降低梯度方差、提高硬件吞吐,但单位 Token 的参数更新次数减少,超过临界 Batch 后收益递减,甚至降低泛化或造成学习率不匹配。大模型通常按全局 Token 数定义 Batch,通过数据并行和梯度累积实现;扩 Batch 时要联动学习率、warmup、总训练 Token 和稳定性,而不是只把显存塞满。

详细版

全局 Batch 近似为 micro_batch × gradient_accumulation × data_parallel_world_size,对变长序列更准确的口径是每次更新的有效 Token。若 64 个数据并行 rank,每卡 micro-batch 2,累积 8 次,则每次更新含 1024 条序列;序列长 2048 时约 210 万 Token。

小 Batch 梯度噪声大但更新频繁,大 Batch 方向稳定并提高矩阵计算效率。学习率可从线性或平方根缩放作为起点,但需实验校准;保持总 Token 不变时,Batch 翻倍意味着 optimizer step 减半,warmup 应优先按 Token 比例定义。应扫描吞吐、loss-per-token、梯度范数和验证任务,寻找临界点而非最大可运行值。

每次更新 Token = micro_batch × accumulation × DP × 有效序列长度
更新次数       = 总训练 Token / 每次更新 Token

完整版教学

一、Batch 改变的是梯度估计

真实目标梯度是整个数据分布的期望,mini-batch 只提供随机估计。Batch 越大,样本噪声通常越小,更新方向更稳定;Batch 越小,噪声更强,可能帮助探索但也容易出现尖峰。它影响的不只是吞吐,还改变优化动力学。

如果固定训练 1000 亿 Token,把每步 Batch 从 100 万 Token 增到 400 万 Token,更新次数会从 10 万降到 2.5 万。即使每步梯度更准,模型也获得了更少的参数更新机会,所以不能只比较相同 step 下的 loss。

二、先统一 Batch 的口径

“Batch=1024”可能指样本、序列或 Token,变长训练中差别很大。推荐记录 global batch tokens,并同时披露 micro-batch、梯度累积和数据并行规模:

B_global_seq = B_micro × A × D
B_global_tok = Σ non_padding_tokens over one optimizer update

例如 B_micro=2、累积 A=8、DP=64,得到 1024 条序列。若平均有效长度仅 1500 而非上限 2048,真实 Batch 是 153.6 万 Token,不是 209.7 万;padding 不应计入学习数据量。

三、临界 Batch 与收益递减

在小 Batch 区间,增大 Batch 能显著降低噪声,常可近似按比例提高学习率并减少步数。达到临界 Batch 后,相邻样本带来的梯度信息高度重复,再增大 Batch 主要换取并行吞吐,优化步数减少却无法等比例补偿。临界点会随训练阶段、模型与数据变化。

记忆钩子:Batch 扩大先是在“减少无用噪声”,过了临界点就变成“用更多样本换同一次更新”,边际收益自然下降。

Gradient Noise Scale 可用于估计这个区域,但最终仍需 loss-per-token 和下游评测验证。训练早期与后期的合适 Batch 也可能不同,因此有些方案采用 Batch ramp-up。

四、学习率为什么必须联动

大 Batch 梯度平均更稳定,通常能承受更大学习率。常见启发式有线性缩放 lr ∝ B 和平方根缩放 lr ∝ sqrt(B),但 Adam、归一化、模型深度都会改变规律。直接线性放大可能造成 loss spike,尤其在 warmup 不足时。

变化需要一起检查
Batch ×2LR 候选、更新次数
DP ×2全局 Batch 是否意外翻倍
累积步数 ×2通信频率、梯度缩放
序列长度 ×2Token Batch、激活显存
总 Token 固定Scheduler 应按 Token 对齐

更可靠做法是小规模扫描若干 LR,并以相同已见 Token 比较,而不是照搬别的模型经验。

五、梯度累积等价吗

在没有 BatchNorm、随机性可控且对所有 micro-batch 累加后只更新一次时,梯度累积在数学上接近大 Batch。但 Dropout 随机流、梯度裁剪位置、混合精度溢出和分布式通信实现会造成差异。梯度应先正确按总样本或 Token 归一化,再裁剪和 optimizer step。

累积可降低单卡激活显存,却不会减少每个样本的前后向计算,还降低 optimizer 更新频率。通信可用 no_sync 在中间 micro-step 跳过 all-reduce,但最后一步必须同步。发生 NaN 时要明确丢弃整个累计窗口还是局部 batch。

六、吞吐与显存如何权衡

Micro-batch 太小会让矩阵形状和 GPU 利用率不佳,太大则激活 OOM。通常先在单卡上找能高效运行的最大 micro-batch,再用累积达到优化目标 Batch,最后扩 DP。吞吐应按有效 Token/s 测量,并排除 padding 与数据等待。

大 Batch 还会增加数据读取突发和 all-reduce 压力。若扩卡后 global batch 不变,每卡工作变少,通信占比可能上升;若 global batch 同步扩大,则优化超参数也变了。性能实验与收敛实验必须明确区分强扩展和弱扩展。

七、如何设计 Batch 扫描实验

选 4~6 个 Batch 候选,保持数据顺序、总 Token、模型和评测一致,为每个候选调一小组学习率。记录训练 loss 对已见 Token、验证 loss、梯度范数、溢出次数、Token/s 和单位 Token 成本。最终选择满足质量的最高吞吐点,而不是单纯最大 Batch。

例如 Batch 为 1M、2M、4M Token 时吞吐分别为 100、170、220 万 Token/s,但达到同一 loss 所需 Token 为 80B、82B、105B。4M 虽瞬时最快,数据效率明显变差;综合墙钟时间与计算成本,2M 可能更优。

八、常见误区与追问

  • 误区:Batch 越大,收敛一定越快。 每步更稳定不等于按 Token 或墙钟更快,超过临界点收益递减。
  • 误区:增加 GPU 数但不改配置,优化过程不变。 若每卡 Batch 不变,全局 Batch 会随 DP 放大。
  • 误区:梯度累积只影响显存。 它还改变 optimizer 更新与通信节奏,错误归一化会改变有效学习率。
  • 追问:warmup 按 step 还是 Token? 跨 Batch 实验更适合按已见 Token 或训练比例定义。
  • 追问:为什么变长数据要用 Token Batch? 相同序列数可能包含完全不同的有效训练量和计算量。
  • 追问:如何判断 Batch 过大? 增大后吞吐收益有限、达到同一 loss 所需 Token 上升,且任务质量不再改善。

九、加强记忆

Batch Size 可记成“口径、噪声、步数、联调、拐点”:先用全局有效 Token 统一口径,大 Batch 降低梯度噪声却减少更新步数,学习率与 warmup 必须联调;再结合累积和 DP 找硬件效率,最后用相同总 Token 的质量—吞吐曲线找临界拐点。能把优化效率与系统吞吐分开,就回答到了本质。