梯度累积与梯度检查点
两种常见显存优化手段:梯度累积用多个 micro-batch 合成一次参数更新;梯度检查点少保存激活、反向时重新计算。前者换更新频率,后者换计算时间。
梯度累积(Gradient Accumulation)
Section titled “梯度累积(Gradient Accumulation)”若每张设备的 micro-batch 为 ,累积步数为 ,设备数为 ,则常用的有效 batch size 为:
手写训练循环时,如果每个 micro-batch 的 loss 默认取 mean,常需除以 ,并且只在累积边界调用 optimizer.step():
optimizer.zero_grad()for i, batch in enumerate(loader): loss = model(**batch).loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()Trainer 一类框架通常会处理 loss 缩放、同步和最后一个不满周期的 batch;不能机械重复手写逻辑。
累积允许用较小 micro-batch 完成较大的有效 batch,但不会消除模型权重、优化器状态或单个 micro-batch 激活。一次 optimizer update 需要更多前后向步骤,吞吐和调参语义也会改变。
梯度检查点(Gradient Checkpointing)
Section titled “梯度检查点(Gradient Checkpointing)”普通反向传播保存大量中间激活。梯度检查点只保存部分边界张量,反向经过被检查点包围的区段时重新执行前向:
更少激活显存 ↔ 更多重计算时间它保存的是 activation checkpoint,不是训练断点文件,也不等于 save_checkpoint()。
两者如何组合
Section titled “两者如何组合”- accumulation 控制一次参数更新由多少 micro-batch 组成;
- checkpointing 控制单个 micro-batch 前向要保留多少激活;
- AMP 改变部分算子的 dtype;
- LoRA 减少可训练参数与优化器状态。
四者可组合,但每一种优化影响的显存组成不同,不能简单把节省比例相乘。
课程证据与边界
Section titled “课程证据与边界”第 11 讲给出 Trainer 的 gradient_accumulation_steps 和显存观察;第 13 讲展示 checkpointing 的激活重算及保存输出;第 14 讲在 Llama 2 SFT 中组合这些配置。
课程实验的对照有时同时改变 micro-batch、累积步数或其他设置,因此具体 MB 数不应当作严格消融结论。概念层只保留机制与可复核公式。
- 累积步数翻倍,显存必减半:单个 micro-batch 未变时,峰值激活可能几乎不变。
- 有效 batch 相同,训练必完全等价:随机层、归一化、调度器与分布式同步会造成差异。
- checkpointing 不影响速度:它以额外前向重算换显存。