梯度累积


文档摘要

梯度累积 本节摘要:用你买不起的有效批,一次一个微批(micro-batch)来训练。缩放损失、按住优化器步、让梯度在参数缓冲里堆积——这就是梯度累积。等式是 :每个微批的损失除以 后反向,PyTorch 默认把梯度累加进 ,累积完后步一次优化器,产出的梯度缓冲与一次全批反向产出的张量相同(差仅浮点求和顺序,断言 )。优化器状态(动量、Adam 矩)每有效步进一次而非每微批,否则指数滑动均值看到错误频率、烧穿调度。单机上是簿记;多 rank 集群上,非末微批包进 跳过梯度 all-reduce,末微批一次 reduce 全累积梯度,省 N 倍网络成本。 对应原课程:Phase 19 · Lesson 46 · (原英文 )。本节属「预训练/分布式」赛道第五节。

梯度累积

本节摘要:用你买不起的有效批,一次一个微批(micro-batch)来训练。缩放损失、按住优化器步、让梯度在参数缓冲里堆积——这就是梯度累积。等式是 effective_batch = micro_batch * accum_steps:每个微批的损失除以 accum_steps 后反向,PyTorch 默认把梯度累加进 param.grad,累积完后步一次优化器,产出的梯度缓冲与一次全批反向产出的张量相同(差仅浮点求和顺序,断言 max-abs-diff < 1e-4)。优化器状态(动量、Adam 矩)每有效步进一次而非每微批,否则指数滑动均值看到错误频率、烧穿调度。单机上是簿记;多 rank 集群上,非末微批包进 no_sync 跳过梯度 all-reduce,末微批一次 reduce 全累积梯度,省 N 倍网络成本。

对应原课程:Phase 19 · Lesson 46 · gradient-accumulation(原英文 phases/19-capstone-projects/46-gradient-accumulation/docs/en.md)。本节属「预训练/分布式」赛道第五节。

学习目标

阅读完本节,你应当能够:

  1. 推导有效批等式:effective_batch = micro_batch * accum_steps
  2. 实现每微批损失缩放,使累积梯度匹配单次全批反向。
  3. 跳过优化器同步直到末微批(sync-on-last-step)。
  4. 读吞吐对有效批的曲线,解释收益递减。

一、问题与直觉

你想用有效批 512 训练,因损失曲线更平滑、优化器步在那个尺度更合理。桌上的加速器装 32 例就内存爆。翻倍批不行,砍半模型也不行。2017 年以来整个领域用的招是:跑 16 次反向,让梯度在参数缓冲里累积,只在计数到目标时步优化器。

风险是损失不再是更大批时的同一个数。16 个微批的交叉熵朴素求和是单全批损失的 16 倍。不缩放,梯度方向对但量级错,优化器步大 16 倍。修复是一个除法——也容易忘。

契约很短:每微批损失除以 accum_stepsbackward()(PyTorch 默认把梯度求和进 param.grad,除法把运行和推回正确尺度);优化器步每有效批一次,在末微批反向后(中途步会让运行余下参数偏);优化器状态每有效步进一次而非每微批;多 rank 上非末微包包进 no_sync 跳过 all-reduce,末微批一次 reduce。

二、从零实现

代码里的等价性证明

# 单全批 loss = criterion(model(x_full), y_full) loss.backward() opt.step()

等价于(差仅浮点求和顺序):

for x, y in chunks(x_full, y_full, n): scaled = criterion(model(x), y) / n # 缩放损失 scaled.backward() opt.step()

末尾的累积梯度缓冲与单全批反向产出的张量相同。本节代码在 equivalence_check 里断言 max-abs-diff < 1e-4

sync-on-last-step 模式

def train_one_optimizer_step(model, opt, micro_batches, accum_steps): opt.zero_grad() for i, (x, y) in enumerate(micro_batches): is_last = (i == accum_steps - 1) ctx = nullcontext() if is_last else no_sync_context(model) # 非末跳 all-reduce with ctx: loss = criterion(model(x), y) / accum_steps # 缩放 loss.backward() # 离开末微批的 no_sync,梯度在多 rank 上已 all-reduce torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step() # 每有效步一次

code/main.py 做三件事:equivalence_check() 建两个同种子网络副本,一个看 16 样本单前向,一个看四个 4 样本块(损失除四),比步前梯度缓冲与步后参数;train_one_optimizer_step 走微批,非末进 no_sync_context(单进程是 no-op,DDP 上跳 all-reduce),sync_counter 记离开 no_sync 的次数,对 N 微批每有效步计一而非 N;sweep_effective_batches 用固定微批与一组累积步跑同模型,每配置日志 samples_per_secmedian_step_mssync_callsavg_loss,落 outputs/accum-curve.json

设计要点:成本无免费午餐——翻倍 accum_steps 翻倍每优化器步的墙时间,变的是梯度估计的方差:同墙预算下优化器步更少但每步平均更多样本。大批与小批在文献里被当不同优化问题;本节的机理是机械的,非统计的。no_sync 在单机是 no-op 但保留调用点,使代码与 DDP 循环逐字相同——迁移性靠它。

三、框架对比

HuggingFace Trainergradient_accumulation_steps 参数把这一切打包:配 gradient_accumulation_steps=16,内部跑微批、缩放、末微批同步。PyTorch FSDP/DDP 原生支持 model.no_sync() 上下文。DeepSpeed、Megatron-LM 在分布式上自管累积与梯度分桶通信。本节手写让你看清等价性证明、缩放的那个除法、sync-on-last-step 的簿记。生产上累积步常配数据并行(每 rank 各自累积,末微批跨 rank all-reduce)与流水并行(不同层在不同 rank,累积是每 rank 局部的),但累积的核心契约与本节一致。

四、可复用产物

code/main.py + outputs/accum-curve.json:equivalence_checktrain_one_optimizer_stepsweep_effective_batches 均可复用。demo 打印等价性 diff、扫描表、JSON 路径,退零。train_one_optimizer_step 的 sync-on-last-step 模式是生产累积步的标准骨架——换到 DDP/FSDP 上,只需把 no_sync_context 换成 model.no_sync(),其余不动。

五、练习

  1. 等价性断言:给 equivalence_check 加更多种子与批大小,确认 max-abs-diff 始终 < 1e-4。
  2. 忘缩放:故意去掉损失除法,确认梯度量级大 accum_steps 倍、训练发散。
  3. 中途步:在非末微批步优化器,确认参数偏移、后续梯度基于错位参数。
  4. 吞吐曲线:扫 accum_steps 从 1 到 32,绘 samples_per_secmedian_step_ms,确认前者饱和、后者线性增长。
  5. DDP no_sync:把 no_sync_context 接成 model.no_sync()(若有双进程环境),确认 sync_calls 每有效步为一。

本节要点回顾

  1. 等式:effective_batch = micro_batch * accum_steps,用累积换内存。
  2. 缩放的那个除法:每微批损失除 accum_steps,否则梯度量级大 N 倍。
  3. 步每有效批一次:优化器状态每有效步进,否则动量看到错误频率。
  4. 等价性可证:累积梯度缓冲与单全批反向相同,断言 max-abs-diff < 1e-4
  5. sync-on-last-step:非末微批 no_sync 跳 all-reduce,末微批一次 reduce。
  6. 无免费午餐:翻倍累积翻倍墙时间,变的是梯度方差而非吞吐。

下一节,我们做「检查点保存与恢复」——原子写五桶状态(模型/优化器/调度器/计数器/RNG),使任意时刻的崩溃留下有效文件。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U