大 Batch 放不进显存怎么办?梯度累积的精确归一化与 DDP no_sync
从梯度平均的分母出发,手算不等长 micro-batch 的累积误差,实现 token 精确归一化、AMP 与 DDP no_sync,并覆盖尾批、裁剪、调度和等价性调试。
上一篇用 Activation Checkpointing(激活检查点)重算中间激活,换回了部分显存。但单次前向能容纳的样本或 token 仍有上限。最直接的办法是把一个大 batch 拆成多个 Micro-batch(微批次),依次反向,把梯度留在参数的 .grad 中,最后只更新一次。
本篇聚焦一个容易被“除以累积步数”掩盖的问题:Gradient Accumulation(梯度累积)何时真的等价于一个大 batch,以及可变 token 数与 DDP 下怎样得到正确分母。
01 梯度为什么能够相加?#
一个更新窗口有 个有效监督 token,每个 token 的 loss 为 。目标是
微批次 包含有效 token 集 ,先对每批求 loss sum:
只要模型参数在窗口内不更新,反向的加法就与把这些 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 时钟]mermaid02 “每批 mean 再除以 K”为什么会错?#
设两个 micro-batch 分别有 2 和 6 个有效 token,各自 token loss 为:
micro 1: [2, 4] mean = 3
micro 2: [1, 1, 1, 1, 1, 1] mean = 1text若计算 (3 + 1) / 2 = 2,两个 micro-batch 权重相同;真正的大 batch 平均是
只有每个 micro-batch 的有效元素数相等时,“各自 mean 再除以 ”才成立。语言模型有 padding、文档边界 mask 和不同长度,必须累加 loss sum 与有效 token 总数。
03 单卡上的精确实现#
torch.nn.functional.cross_entropy(..., reduction="sum", ignore_index=-100) 会让被忽略位置不贡献 loss。下面的 loss_sum.backward() 将未归一化梯度直接相加,窗口结束后统一除以有效 token 数。
import torch
import torch.nn.functional as F
optimizer.zero_grad(set_to_none=True)
window_tokens = 0
for micro in micro_batches:
input_ids = micro["input_ids"].cuda() # [B_k,L_k]
labels = micro["labels"].cuda() # [B_k,L_k], pad=-100
logits = model(input_ids) # [B_k,L_k,V]
loss_sum = F.cross_entropy(
logits.transpose(1, 2), labels,
ignore_index=-100,
reduction="sum",
) # scalar
loss_sum.backward() # 参数 .grad 累加 sum
window_tokens += int(labels.ne(-100).sum())
if window_tokens == 0:
raise RuntimeError("accumulation window has no supervised token")
for parameter in model.parameters():
if parameter.grad is not None:
parameter.grad.div_(window_tokens)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
optimizer.zero_grad(set_to_none=True)python输入为整数 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 的梯度是在新参数 上计算;目标变成两次小 batch SGD,不再是同一个 上的大 batch 梯度。
一个窗口内,以下操作都只能在最后执行一次:
- 梯度归一化与 clipping;
optimizer.step();optimizer.zero_grad();- AMP scaler 的
unscale_、step与update; - 按成功更新计数的学习率与 token 时钟。
05 与 AMP 组合时正确顺序是什么?#
上一篇已说明同一窗口必须保持同一 scale。为了让动态 GradScaler(梯度缩放器)检查正确的梯度,先让所有 micro-batch 用 scaler.scale(loss_sum).backward(),窗口末尾再 unscale,然后除以全局 token 分母、裁剪和更新。
optimizer.zero_grad(set_to_none=True)
window_tokens = 0
for micro in micro_batches:
with torch.autocast("cuda", dtype=torch.float16):
logits = model(micro["input_ids"])
loss_sum = F.cross_entropy(
logits.transpose(1, 2), micro["labels"],
ignore_index=-100, reduction="sum",
)
scaler.scale(loss_sum).backward()
window_tokens += int(micro["labels"].ne(-100).sum())
scaler.unscale_(optimizer) # 每窗口一次
for p in model.parameters():
if p.grad is not None:
p.grad.div_(window_tokens) # 现在已是 unscaled grad
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()python先除 token 分母再裁剪,因为阈值通常定义在平均梯度上。若先裁剪 loss sum 梯度,窗口 token 数翻倍就会凭空改变裁剪强度。
06 DDP 为什么会让中间 micro-batch 白白通信?#
DistributedDataParallel(分布式数据并行,DDP)默认在 backward 中对梯度 bucket 做 all-reduce。若每个 micro-batch 都同步,通信发生 次;但更新只需要最终累积和同步一次。
PyTorch 2.14 的 ddp.no_sync() 会暂缓梯度同步,第一次离开该上下文的 forward-backward 再同步累积梯度。官方特别提醒:forward 也必须放进 no_sync() 上下文,否则仍会同步。
from contextlib import nullcontext
optimizer.zero_grad(set_to_none=True)
local_tokens = 0
for i, micro in enumerate(micro_batches):
is_last = i == len(micro_batches) - 1
sync_context = nullcontext() if is_last else ddp.no_sync()
with sync_context: # 包住 forward + backward
logits = ddp(micro["input_ids"])
loss_sum = F.cross_entropy(
logits.transpose(1, 2), micro["labels"],
ignore_index=-100, reduction="sum",
)
loss_sum.backward()
local_tokens += int(micro["labels"].ne(-100).sum())python最后一个 backward 会把此前本地累积的梯度一起同步。若最后一次也用了 no_sync(),各 rank 会拿不同梯度继续更新,模型副本从此分叉。
07 全局 token 分母为何还要乘 world size?#
rank 的本地梯度和为 。DDP 同步后参数 .grad 是
全局平均目标应为
因此先 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(余批)。三种选择要显式定义:
- 照常更新:用实际
global_tokens归一化;优化步的 batch 较小,但不丢数据。 - 跨 epoch 延续:保留梯度和计数到下一轮;数据顺序与 checkpoint 恢复更复杂。
- 丢弃余批:复现简单,但每轮系统性丢样本,分布式各 rank 必须一致。
绝不能仍除以配置的 。恢复 checkpoint 时若允许保存“半个窗口”,必须同时保存已累积梯度、micro-step、token 计数、scaler 与数据游标;工程上更常在更新边界保存。
09 有效 batch 大小应该怎样描述?#
固定形状视觉任务常写
但语言模型更应报告每次 update 的全局有效 token:
它同时反映 padding、loss mask、sequence packing 和 rank 间长度差异。日志至少保存 micro_step、optimizer_step、local/global_effective_tokens、physical_tokens、grad_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 越大越好?#
更大的 降低每个 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 今天真正需要记住什么?#
- 梯度累积等价于大 batch 的前提,是在同一参数点累加 loss sum,最后除以全局有效元素数。
- 可变长度任务不能用 mean-of-means;要显式记录 loss sum 与有效 token。
- DDP 中非最后 micro-batch 用
no_sync(),且上下文必须同时包住 forward 和 backward。 - DDP 默认平均梯度,所以本地 sum loss 的最终缩放是
world_size / global_tokens。
16 思考题与小练习#
- 三个 micro-batch 的有效 token 数为
[3,5,2],mean loss 为[2,1,4]。分别计算错误的 mean-of-means 与正确全局均值。 - 写一个两进程 DDP 小测试,让两个 rank 分别拥有 2 和 6 个 token,验证缩放因子
world_size/global_tokens与单进程 8-token 梯度一致。 - 为“尾窗口照常更新”设计 checkpoint 与日志字段,保证中断恢复不会重复或漏掉样本。
相关工作#
- Goyal et al., Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour ↗,讨论大 batch、学习率缩放与 warmup 的经验规律。
- McCandlish et al., An Empirical Model of Large-Batch Training ↗,用梯度噪声尺度分析 batch 增大何时仍有效。
- Ott et al., Scaling Neural Machine Translation ↗,展示梯度累积和大 batch 在神经机器翻译训练中的作用。
- Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training ↗,解释 DDP 的梯度 bucket、同步与工程设计。
- Smith et al., Don’t Decay the Learning Rate, Increase the Batch Size ↗,比较学习率衰减与逐步增大 batch 的关系。
17 下一篇预告#
单卡和数据并行的 batch 语义已经清楚,下一篇将继续研究大模型如何跨设备放置参数与 optimizer state,并比较 Data Parallel、Tensor Parallel 与 Pipeline Parallel 各自在切什么。