一篇读懂梯度累积:显存不够,步数来凑

微调 7B 模型时最常见的窘境:教程说 batch size 该开 64,你的显存只装得下 4。把 batch 调小硬训?训练会不稳,效果打折(学习率与 batch size 是联动的,见《一篇读懂学习率》)。**梯度累积(gradient accumulation)**就是解这道题的标准工具:显存一次只装 4 条数据,但梯度攒 16 个微批再更新一次——等效 batch size 照样是 64。

原理:把一次大更新拆成多次小记账

先回忆训练一步在做什么:取一个 batch → 前向算 loss → 反向算每个参数的梯度 → 优化器按梯度更新参数。关键观察是:梯度是各样本 loss 梯度的平均值,而平均值可以分批累加。

64 条数据的梯度,等于 16 个「4 条数据微批」的梯度之和除以 16。所以只要在微批之间不做参数更新、也不清空梯度,把 16 份梯度攒齐后再除以步数、更新一次,效果在数学上与「一次吃 64 条」几乎等价(例外见下文 BatchNorm)。

import torch

model.zero_grad()                      # 开局清零梯度
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5)

for step, micro_batch in enumerate(loader):
    loss = model(micro_batch).loss / ACC_STEPS   # ① loss 先除以累积步数
    loss.backward()                              # ② 梯度累加到 .grad(不清零)

    if (step + 1) % ACC_STEPS == 0:              # ③ 攒够了才真正更新
        optimizer.step()
        optimizer.zero_grad()

三个标注点是全部精髓:① loss 除以累积步数,保证等效梯度和「真大 batch」一样是平均而非求和;② backward() 默认把梯度累加进 .grad,这正是「不清零就能攒」的原因;③ 攒满才 step(),每个累积周期只更新一次。等效 batch size 的公式:微批 × 累积步数 × 卡数。

陷阱一:BatchNorm 让「等效」打折扣

梯度累积「完全等价于大 batch」有一个著名例外:BatchNorm。BN 的均值方差是按当前喂进去的微批算的——也就是说,网络里 BN 层看到的有效 batch 永远是微批大小(4),而不是等效的 64。fast.ai 论坛的说法很准确:你实际上有两个 batch size,BN 之外的一切等效于 64,BN 只等效于 4。微批很小时,BN 统计量噪声大,训练质量随之劣化。

好在Transformer 系模型(现代 LLM 的全部)用的是 LayerNorm/RMSNorm——归一化在样本内部统计,不跨样本、不依赖 batch(见《一篇读懂归一化》),所以微调 LLM 几乎不受此坑影响。但训练 CNN 或微调视觉模型时要留意,缓解办法:换 GroupNorm/LayerNorm、微批别太小、或用跨卡同步 BN。

陷阱二:混合精度的缩放器要配合

用混合精度训练(见《一篇读懂混合精度》)时,GradScaler 的使用方式要变:scaler 在整个累积周期只 step 一次,而且要在 scaler.unscale_(optimizer) 之后才能做梯度裁剪。一个常见的翻车写法是每个微批都 scaler.step()——梯度没攒齐就更新,等于没累积,还可能触发频繁的降缩放。标准写法:

if (step + 1) % ACC_STEPS == 0:
    scaler.unscale_(optimizer)     # 先反缩放,才能做 grad clip
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    scaler.step(optimizer)         # 每周期只 step 一次
    scaler.update()
    optimizer.zero_grad()

陷阱三:零头与归一化口径

数据集长度不能被「微批 × 累积步数」整除时,最后一个累积周期不满额——损失归一化口径就和真大 batch 有细微出入。多数框架(HuggingFace Trainer 等)已处理,但要意识到这个边角的存在;自己手写循环时,可以按「实际攒到的样本数」归一化,或容忍最后一个不完整周期。

另一个常见困惑:累积期间的梯度放在哪?答案是显存里的 .grad 缓冲——它与激活值显存是两码事。所以梯度累积省的是激活值显存(微批变小),参数与梯度/优化器状态的显存一点没省(这部分要省得上 LoRA 或量化优化器,见《一篇读懂 LoRA》)。

怎么选参数

  • 先定目标等效 batch:按任务与学习率惯例选(LLM 预训练常见百万 token 级 batch,微调常见 16–128 条),再反推「微批 × 步数」组合;
  • 微批尽量吃满显存:微批越大吞吐越高(并行度高),累积步数相应减少——「16 × 4」通常优于「4 × 16」;
  • 调学习率按等效 batch 走:改累积步数相当于改 batch size,学习率要联动(线性缩放或按经验微调),只改一个不改另一个是常见翻车源。

结语

梯度累积是「用时间换显存」的交易:同样的等效 batch,通信与更新次数变少(分布式下还有带宽红利),代价是每轮参数更新的等待变长。理解它只需要记住一句话——梯度是平均值,平均值可以攒;而用好它,则需要躲开 BatchNorm 的统计口径与 AMP 缩放器的两个坑。这三件事想明白了,你就再也不会被「显存 OOM 但又想用大 batch」卡住了。

参考资料

← 返回资讯列表

读者留言

COMMENTS 暂无
仅本站原创文章开放留言 · 请勿留下手机号、邮箱等个人信息

还没有留言,来说第一句?