梯度累积解决什么问题?和增大 Batch Size 一样吗?
简化版
梯度累积把多个 micro-batch 的梯度相加后再执行一次 optimizer step,用较小显存模拟更大的优化 batch。若损失按累积步数正确缩放,结果接近大 batch;但 BN 仍逐 micro-batch 统计,随机层和调度步数也可能造成差异。
详细版
-
每个 micro-batch 前向和反向后不清梯度,累计 K 次再更新。
-
若 loss 默认取 batch mean,通常每次应除以 K,避免梯度放大 K 倍。
-
最后不足 K 个 micro-batch 时要按实际数量处理。
-
梯度裁剪应在 AMP unscale 后、optimizer step 前对累计梯度执行一次。
-
学习率调度器通常按 optimizer update 而非 micro-step 前进。
完整版教学
一、它扩大优化批次,不扩大单次前向批次
参数在 K 次反向之间保持不变,因此累计的是同一点上各 micro-batch 的梯度和。
对可加损失,这与拼成大 batch 的平均梯度接近。
“接近”而非绝对相同,因为 BN 的统计、Dropout 随机掩码、浮点求和顺序和分布式通信时机会不同。
二、底层机制与公式
g = (1/K) * sum_(i=1..K) grad loss_i
optimizer.step() once after K backward calls
effective_batch = micro_batch * K * world_size
三、带数字的推演
每卡 micro-batch=8,累积 K=4,使用 2 张卡,则有效优化 batch 为 8×4×2=64。
若每次 mean loss 不除以 4,更新梯度会放大约 4 倍。
四、方案对比
| 方案/对象 | 核心特点 | 代价或边界 |
|---|---|---|
| 真实大 batch | 一次计算完整统计 | 显存高 |
| 梯度累积 | 省激活显存 | 训练时间与状态处理更复杂 |
| 梯度检查点 | 重算激活省显存 | 不改变 batch 语义 |
五、执行流程
zero_grad -> 重复 K 次 forward/backward(loss/K)
-> AMP unscale -> clip accumulated grads -> optimizer.step -> scheduler.step
六、边界条件与工程代价
在 DDP 中每个 micro-step 都 all-reduce 会浪费通信,可在非更新步使用 no_sync,最后一步再同步。
实现错误会导致不同 rank 梯度不一致。
累积不能修复小 batch BN;要扩大统计 batch 需 SyncBN、冻结 BN 或换 GN。
优化 batch 与归一化统计 batch 必须分开描述。
记忆钩子:牢牢记住两套时钟:micro-step 负责攒梯度,update-step 才改参数、推进调度。
七、常见误区与追问
-
误区:累积 4 次会让 BN 看到 4 倍样本。 BN 每次前向独立统计。
-
追问:为什么 loss 常除以 K? 让累计梯度等于 K 个批次的平均而不是总和。
-
误区:每个 micro-step 都调用 scheduler.step。 多数调度按参数更新次数定义。
-
追问:裁剪在什么时候做? 全部梯度累计并完成 AMP unscale 后、更新前。
-
追问:最后不足 K 步怎么办? 按实际累积数归一化并执行一次更新,不能静默丢弃。
八、加强记忆
牢牢记住两套时钟:micro-step 负责攒梯度,update-step 才改参数、推进调度。
有效 batch 变大,但 BN 统计 batch 没变。