梯度累积:小显存跑大批量—微批次加权求和一步更新
梯度累积通过多个微批次求和后再更新参数,在显存不足时仍能达到大批量的训练效果。
梯度累积通过多个微批次求和后再更新参数,在显存不足时仍能达到大批量的训练效果。
这个机制是免费的。PyTorch 的 .backward() 累加到每个参数的 .grad 里;它永远不会覆盖。通常你通过每步调用 zero_grad() 隐藏这一点。积累就是……不这样做——它让多次后向传播堆积到同一个缓冲区。
loss_a.backward() # p.grad = g_a
loss_b.backward() # p.grad = g_a + g_b <-- ADDED, not replaced
# skip zero_grad() and gradients accumulate across micro-batches
输入 N 个大小为 m 的微批次,对每个执行后向传播,但只每 N 步调用一次 optimizer.step()。微批次之间缓冲区增长;只有一个微批次的激活(m)在活跃,所以峰值激活内存保持平坦——与 N 无关。
for i, (xb, yb) in enumerate(loader): # micro-batches of size m
criterion(model(xb), yb).backward() # add into .grad
if (i + 1) % N == 0: # every N micro-batches...
opt.step(); opt.zero_grad() # ...one update, then reset
每个微批次损失已经是 m 上的均值。把 N 个加起来得到 N× 的 m·N 上的均值——大了 N 倍。把每个损失除以 N,累积缓冲区就变成了有效批次上的真实均值,与真实大批次 B = m·N 位对位地匹配,因为等大小均值的均值就是总体均值。忘记它,你就无声地把学习率乘以了 N。
loss = criterion(model(xb), yb) / N # <-- scale so the SUM is a MEAN
# without /N: buffer = Σ mean_j = N × true mean (the #1 accumulation bug)
# with /N: buffer = Σ mean_j / N = the big-batch mean ✓
总 FLOPs 不变——无论哪种方式你都处理相同的 B 个样本——所以你花费的是时间:每次更新 N 个按顺序的更小的通道,更新频率降低 N 倍。这与梯度检查点(Day-48)形成了鲜明对比,后者增加了一个重计算通路。检查点在固定批大小下用计算换取激活内存;积累在固定内存下用步骤/时间换取批大小——两者清晰地组合。
把它们组合起来,就是围绕你正常循环的五行——与 m·N 大批次相同的权重,内存只需要 m。它与混合精度(autocast + GradScaler)以及梯度检查点清晰地配对。
opt.zero_grad()
for i, (xb, yb) in enumerate(loader): # micro-batches, size m
loss = criterion(model(xb), yb) / N # scale so the sum is a mean
loss.backward() # accumulate into .grad
if (i + 1) % N == 0:
opt.step(); opt.zero_grad() # one real update every N
等价性只对每样本损失成立。BatchNorm 在当前微批次 m 上正则化,而不是有效的 m·N,所以有 BN 时积累不等同于真实大批次。改用 GroupNorm、LayerNorm 或 SyncBN。也要注意尾部——如果加载器不是 N 的倍数,在循环后刷新剩余组,否则你会无声地丢弃这些样本。
大批次是一个内存问题,不是数学问题。.grad 已经求和,所以给它 N 个微批次,记住 1/N,步进一次——你就以一个小批次的代价买到了大批次。