观文听傑

返回

上一篇用 Activation Checkpointing(激活检查点)重算中间激活,换回了部分显存。但单次前向能容纳的样本或 token 仍有上限。最直接的办法是把一个大 batch 拆成多个 Micro-batch(微批次),依次反向,把梯度留在参数的 .grad 中,最后只更新一次。

本篇聚焦一个容易被“除以累积步数”掩盖的问题:Gradient Accumulation(梯度累积)何时真的等价于一个大 batch,以及可变 token 数与 DDP 下怎样得到正确分母。

01 梯度为什么能够相加?#

一个更新窗口有 NN 个有效监督 token,每个 token 的 loss 为 i(θ)\ell_i(\theta)。目标是

L(θ)=1Ni=1Ni(θ),θL=1Ni=1Nθi\mathcal L(\theta)=\frac{1}{N}\sum_{i=1}^{N}\ell_i(\theta),\qquad \nabla_\theta\mathcal L=\frac{1}{N}\sum_{i=1}^{N}\nabla_\theta\ell_i

微批次 kk 包含有效 token 集 SkS_k,先对每批求 loss sum:

gsum=k=1KθiSki,qquadg=gsumkSkg_{\text{sum}}=\sum_{k=1}^{K}\nabla_\theta\sum_{i\in S_k}\ell_i,qquad g=\frac{g_{\text{sum}}}{\sum_k |S_k|}

只要模型参数在窗口内不更新,反向的加法就与把这些 token 拼成一次大 batch 等价;浮点求和顺序、随机层和 BatchNorm 统计仍会造成小差异。

flowchart LR
  A[micro 1: loss_sum + n1] --> G[累加参数 .grad]
  B[micro 2: loss_sum + n2] --> G
  C[micro K: loss_sum + nK] --> G
  A --> N[累加有效 token N]
  B --> N
  C --> N
  G --> D[grad /= global N]
  N --> D
  D --> E[unscale / clip / optimizer.step]
  E --> F[清空 grad 并推进 token 时钟]
mermaid

02 “每批 mean 再除以 K”为什么会错?#

设两个 micro-batch 分别有 2 和 6 个有效 token,各自 token loss 为:

micro 1: [2, 4]             mean = 3
micro 2: [1, 1, 1, 1, 1, 1] mean = 1
text

若计算 (3 + 1) / 2 = 2,两个 micro-batch 权重相同;真正的大 batch 平均是

2+4+1+1+1+1+1+18=1.5\frac{2+4+1+1+1+1+1+1}{8}=1.5

只有每个 micro-batch 的有效元素数相等时,“各自 mean 再除以 KK”才成立。语言模型有 padding、文档边界 mask 和不同长度,必须累加 loss sum 与有效 token 总数。

03 单卡上的精确实现#

torch.nn.functional.cross_entropy(..., reduction="sum", ignore_index=-100) 会让被忽略位置不贡献 loss。下面的 loss_sum.backward() 将未归一化梯度直接相加,窗口结束后统一除以有效 token 数。

输入为整数 token [B_k,L_k],logits 为 [B_k,L_k,V],loss 是标量,参数梯度 shape 与对应参数完全相同。set_to_none=True 避免无意义清零写入,并让“本窗口没有梯度”和“梯度全是 0”更容易区分。

04 为什么 optimizer 不能在中途 step?#

如果 micro 1 反向后就更新参数,micro 2 的梯度是在新参数 θ1\theta_1 上计算;目标变成两次小 batch SGD,不再是同一个 θ0\theta_0 上的大 batch 梯度。

一个窗口内,以下操作都只能在最后执行一次:

  • 梯度归一化与 clipping;
  • optimizer.step()
  • optimizer.zero_grad()
  • AMP scaler 的 unscale_stepupdate
  • 按成功更新计数的学习率与 token 时钟。

05 与 AMP 组合时正确顺序是什么?#

上一篇已说明同一窗口必须保持同一 scale。为了让动态 GradScaler(梯度缩放器)检查正确的梯度,先让所有 micro-batch 用 scaler.scale(loss_sum).backward(),窗口末尾再 unscale,然后除以全局 token 分母、裁剪和更新。

先除 token 分母再裁剪,因为阈值通常定义在平均梯度上。若先裁剪 loss sum 梯度,窗口 token 数翻倍就会凭空改变裁剪强度。

06 DDP 为什么会让中间 micro-batch 白白通信?#

DistributedDataParallel(分布式数据并行,DDP)默认在 backward 中对梯度 bucket 做 all-reduce。若每个 micro-batch 都同步,通信发生 KK 次;但更新只需要最终累积和同步一次。

PyTorch 2.14 的 ddp.no_sync() 会暂缓梯度同步,第一次离开该上下文的 forward-backward 再同步累积梯度。官方特别提醒:forward 也必须放进 no_sync() 上下文,否则仍会同步。

最后一个 backward 会把此前本地累积的梯度一起同步。若最后一次也用了 no_sync(),各 rank 会拿不同梯度继续更新,模型副本从此分叉。

07 全局 token 分母为何还要乘 world size?#

rank rr 的本地梯度和为 GrG_r。DDP 同步后参数 .grad

Gddp=1Rr=1RGrG_{\text{ddp}}=\frac{1}{R}\sum_{r=1}^{R}G_r

全局平均目标应为

G=rGrNglobal=GddpRNglobalG=\frac{\sum_r G_r}{N_{\text{global}}} =G_{\text{ddp}}\frac{R}{N_{\text{global}}}

因此先 all-reduce 各 rank 的 local_tokens 得到 global_tokens,再把已同步梯度乘 world_size / global_tokens

count = torch.tensor(local_tokens, device="cuda", dtype=torch.float64)
torch.distributed.all_reduce(count, op=torch.distributed.ReduceOp.SUM)
global_tokens = count.item()
scale = torch.distributed.get_world_size() / global_tokens

for p in ddp.parameters():
    if p.grad is not None:
        p.grad.mul_(scale)
python

这允许各 rank 因长度分桶而有不同有效 token 数,只要它们执行相同数量的 forward-backward 并按同一时刻同步。

08 累积窗口末尾不足 K 批怎么办?#

数据集结束、过滤坏样本或 OOM 重试都可能留下 remainder(余批)。三种选择要显式定义:

  1. 照常更新:用实际 global_tokens 归一化;优化步的 batch 较小,但不丢数据。
  2. 跨 epoch 延续:保留梯度和计数到下一轮;数据顺序与 checkpoint 恢复更复杂。
  3. 丢弃余批:复现简单,但每轮系统性丢样本,分布式各 rank 必须一致。

绝不能仍除以配置的 KK。恢复 checkpoint 时若允许保存“半个窗口”,必须同时保存已累积梯度、micro-step、token 计数、scaler 与数据游标;工程上更常在更新边界保存。

09 有效 batch 大小应该怎样描述?#

固定形状视觉任务常写

Beffective=Bmicro×K×RB_{\text{effective}}=B_{\text{micro}}\times K\times R

但语言模型更应报告每次 update 的全局有效 token:

Neffective=r=1Rk=1KNr,kN_{\text{effective}}=\sum_{r=1}^{R}\sum_{k=1}^{K}N_{r,k}

它同时反映 padding、loss mask、sequence packing 和 rank 间长度差异。日志至少保存 micro_stepoptimizer_steplocal/global_effective_tokensphysical_tokensgrad_norm 与是否成功更新。

10 梯度累积等价性的边界#

组件是否通常等价原因
Linear/LayerNorm近似等价每样本计算不依赖 batch 统计
BatchNorm不等价每个 micro-batch 分别计算均值方差
Dropout统计上接近随机 mask 与大 batch 的调用顺序不同
梯度裁剪可等价必须在累积和归一化后只裁一次
AdamW可等价每窗口只 step 一次,状态只更新一次
学习率日程可等价只在成功 optimizer step 后推进

浮点加法不满足严格结合律,因此即使公式等价也不应要求 bitwise identical(逐位相同)。正确验收是 FP64/FP32 小模型中梯度误差在合理容差内,并且短程 loss 轨迹一致。

11 一个最小等价性测试#

def flatten_grads(model):
    return torch.cat([
        p.grad.detach().flatten()
        for p in model.parameters() if p.grad is not None
    ])

# model_big 与 model_acc 初始 state_dict 完全相同,关闭 dropout
# 路径 A:8 个 token 一次 mean backward
# 路径 B:2 + 6 个 token 分别 sum backward,最后除以 8
g_big = flatten_grads(model_big)
g_acc = flatten_grads(model_acc)
torch.testing.assert_close(g_acc, g_big, rtol=1e-5, atol=1e-7)
python

若失败,按顺序检查:初始参数、样本顺序与 mask、loss reduction、分母、是否中途 zero/step、随机层、BatchNorm,最后才考虑浮点累积顺序。

12 常见错误与最短调试路径#

症状常见原因最短检查
长短 batch 混合后 loss 漂移mean-of-means 权重错误打印每批 loss sum 与有效 token
梯度小了 world size 倍DDP 平均后又直接除全局 token检查 world_size/global_tokens 因子
通信次数没有下降forward 没包在 no_sync()profiler 统计每窗口 all-reduce 数
各 rank 参数逐渐不同最后一个 micro 也禁了同步每次更新后比参数 checksum
clipping 随 K 改变对未归一化 sum gradient 裁剪归一化后再记录 grad norm
恢复后第一步异常在半窗口保存却没恢复 .grad只在 update 边界保存或补齐状态
OOM 后更新权重偏了跳过一批却仍用固定分母从实际成功 micro-batch 重算计数

13 性能上是不是 K 越大越好?#

更大的 KK 降低每个 micro-batch 的激活峰值,并让 DDP 少同步;但它也增加 Python/launch 开销,延迟 optimizer step,并可能让单次矩阵太小而无法吃满 GPU。极大的有效 batch 会降低梯度噪声,未必提升样本效率,学习率也不能无条件线性放大。

应对候选 (micro_batch, K) 组合测 tokens/s、峰值显存、每次 update 时间、通信占比和验证质量。目标是满足显存约束后尽量提高端到端吞吐,而不是最大化累积次数。

14 它会在哪里失败?#

如果模型依赖跨样本操作、批内负样本或 BatchNorm,大 batch 的交互无法由独立 micro-batch 的梯度相加复原。对比学习的分母若需要全局样本,必须先构造正确的跨卡/跨微批负样本集合;否则优化目标已经改变。

梯度累积也不会减少一次 forward 内单个超长样本的激活;那仍需要 sequence parallel、切分 attention、activation checkpointing 或缩短上下文。

15 今天真正需要记住什么?#

  1. 梯度累积等价于大 batch 的前提,是在同一参数点累加 loss sum,最后除以全局有效元素数。
  2. 可变长度任务不能用 mean-of-means;要显式记录 loss sum 与有效 token。
  3. DDP 中非最后 micro-batch 用 no_sync(),且上下文必须同时包住 forward 和 backward。
  4. DDP 默认平均梯度,所以本地 sum loss 的最终缩放是 world_size / global_tokens

16 思考题与小练习#

  1. 三个 micro-batch 的有效 token 数为 [3,5,2],mean loss 为 [2,1,4]。分别计算错误的 mean-of-means 与正确全局均值。
  2. 写一个两进程 DDP 小测试,让两个 rank 分别拥有 2 和 6 个 token,验证缩放因子 world_size/global_tokens 与单进程 8-token 梯度一致。
  3. 为“尾窗口照常更新”设计 checkpoint 与日志字段,保证中断恢复不会重复或漏掉样本。

相关工作#

  1. Goyal et al., Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour,讨论大 batch、学习率缩放与 warmup 的经验规律。
  2. McCandlish et al., An Empirical Model of Large-Batch Training,用梯度噪声尺度分析 batch 增大何时仍有效。
  3. Ott et al., Scaling Neural Machine Translation,展示梯度累积和大 batch 在神经机器翻译训练中的作用。
  4. Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training,解释 DDP 的梯度 bucket、同步与工程设计。
  5. Smith et al., Don’t Decay the Learning Rate, Increase the Batch Size,比较学习率衰减与逐步增大 batch 的关系。

17 下一篇预告#

单卡和数据并行的 batch 语义已经清楚,下一篇将继续研究大模型如何跨设备放置参数与 optimizer state,并比较 Data Parallel、Tensor Parallel 与 Pipeline Parallel 各自在切什么。

大 Batch 放不进显存怎么办?梯度累积的精确归一化与 DDP no_sync
https://zwjcode.cn/blog/gradient-accumulation-token-normalization-ddp-nosync
作者
发布于 2026年9月15日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。