什么是梯度检查点(激活重计算)?它如何用时间换显存?
简化版
梯度检查点(gradient checkpointing,又叫激活重计算) 是一种「用计算换显存」的技巧。反向传播需要前向时的激活值来算梯度,普通做法是把每一层的激活都存下来,显存开销随层数和序列长度线性增长、非常大。梯度检查点的做法是:前向时只保存少数几个”检查点”层的激活,其余丢弃;反向传播到某段时,从最近的检查点重新前向计算出需要的激活。这样激活显存从 O(层数) 降到约 O(√层数),代价是多做一次前向(约增加 30% 计算量)。它是训练长序列/大模型时省激活显存的常用手段。
详细版
为什么激活这么占显存
反向传播算梯度用到链式法则,需要前向时每一层的中间激活。训练时的显存主要由四块组成:参数、梯度、优化器状态、激活。其中激活显存 ∝ batch × 序列长度 × 隐藏维 × 层数,在长序列/大 batch 时会非常庞大,常常是显存瓶颈。
梯度检查点的做法
普通前向:保存每一层的激活 → 反向直接用 → 激活显存 O(L)
检查点: 只保存 √L 个检查点的激活 → 反向时从检查点重算中间激活 → 激活显存 O(√L)
- 把网络分成若干段,只在段边界保存激活(检查点)。
- 反向传播进入某段时,以该段起点的检查点为输入,重新前向一遍算出段内各层激活,再算梯度。
- 用完即丢,不长期占显存。
代价与收益
| 指标 | 普通 | 梯度检查点 |
|---|---|---|
| 激活显存 | O(L) | 约 O(√L) |
| 计算量 | 1 次前向 + 1 次反向 | 约 2 次前向 + 1 次反向(多 ~30%) |
适用场景
- 长上下文训练、大 batch、单卡显存吃紧时。
- 和混合精度、ZeRO、并行等叠加使用。
完整版教学
一、先看清训练显存都花在哪
训练一个模型,显存主要被四部分吃掉:
- 参数:模型权重。
- 梯度:与参数等大。
- 优化器状态:Adam 的动量 m、方差 v、FP32 副本等(往往最大头之一)。
- 激活(activations):前向过程中每一层产生的中间结果,反向要用它们算梯度。
前三块由参数量决定,相对固定;而激活显存随 batch、序列长度、层数线性增长,在长序列训练里会爆炸式增大,常常成为压垮单卡的最后一根稻草。梯度检查点专治这一块。
二、为什么反向传播非要激活不可
反向传播用链式法则逐层往回算梯度。算某一层参数的梯度时,公式里会用到该层前向时的输入/中间激活。举个直觉例子:一个乘法 y = w·x,∂y/∂w = x——要算对 w 的梯度,就必须知道当时的输入 x。所以普通实现会在前向时把所有层的激活缓存下来,供反向使用。层越多、序列越长,这个缓存越大。
三、核心思想:不存就重算
梯度检查点的洞察是:激活是可以”重新算出来”的——只要保留了某个位置的输入,就能通过前向计算再得到后面的激活。既然如此,何必全都存着占显存?不如丢掉大部分,需要时重算。
具体做法(以按段划分为例):
- 把 L 层网络切成若干段,只在段的边界保存激活(这些就是”检查点”)。
- 前向时,段内的中间激活算完就丢弃,不缓存。
- 反向传播到某段时,取该段起点的检查点作为输入,重新前向计算一遍这一段,临时得到段内激活,用来算梯度,算完再丢。
这样常驻的激活显存只剩下「检查点」的量,大幅下降。
四、为什么是 O(√L):一个漂亮的权衡
如果把 L 层均匀分成 √L 段、每段 √L 层:
- 常驻激活 = 检查点数量 ≈ √L 段的边界激活 → O(√L)。
- 重算开销 = 反向时每段重算一次前向,额外约一次完整前向 → 计算量约变成「2 次前向 + 1 次反向」,相比原来的「1 前向 + 1 反向」多约 30%~33%。
这个 √L 的分段是显存与重算成本的平衡点:段分得越多,检查点越少、越省显存,但重算越多。实际框架里可自动或手动选择在哪些层设检查点。
记忆钩子:梯度检查点 = “少存多算”。丢掉中间激活,反向时从最近检查点重算,用约 30% 的额外计算,换来激活显存从 O(L) 降到 O(√L)。
五、在整套显存优化里的位置
梯度检查点只省激活这一块,要和其他手段配合才能训动大模型:
- 梯度检查点:省激活显存(用计算换)。
- 混合精度(BF16):参数/激活/梯度降到 16 位,整体减半。
- ZeRO/FSDP:分片参数/梯度/优化器状态,省前三块的冗余。
- 张量/流水线并行:把模型拆到多卡。
组合拳里,梯度检查点是「针对激活、几乎零副作用(只多点计算)」的常用项,尤其在长上下文训练时几乎必开——因为长序列下激活显存增长最凶。
六、用 16 层分段理解重算
若 16 层网络原本保存每层激活,需要保留 16 份;每 4 层设置一个检查点时,前向只长期保存第 0、4、8、12、16 层附近的边界。反向处理第 13—16 层前,从第 12 层激活重算该段中间值,用完即释放,再处理前一段。
| 策略 | 保存激活 | 额外计算 | 特点 |
|---|---|---|---|
| 不重算 | 约 16 层 | 0 | 显存最高 |
| 每 4 层检查点 | 约 4 个边界 + 一段临时值 | 多次段内前向 | 折中 |
| 更激进重算 | 更少 | 更多 | 速度代价更大 |
forward: save L0 -> recompute L1..L4 -> save L4 -> ...
backward: reload L12 -> recompute L13..L16 -> gradients -> release
重算要求随机操作可复现,否则 dropout 掩码不同会让反向不再对应原前向。验收时应同时比较峰值显存、step time、loss 曲线和梯度一致性,而不是只确认“不 OOM”。
七、常见误区与追问
- 误区:梯度检查点能省参数/优化器显存。 不能,它只省激活显存;参数/梯度/优化器状态要靠 ZeRO、混合精度等。
- 误区:它会降低模型精度。 不会,重算得到的激活和原来完全一致,是精确的,只是多花计算。
- 追问:为什么反向需要激活? 链式法则里参数梯度依赖前向的输入/中间值(如 ∂(wx)/∂w = x)。
- 追问:省了多少、代价多少? 激活显存 O(L)→约 O(√L),计算量约增 30%(多一次前向)。
- 追问:为什么长序列训练一定要它? 激活显存 ∝ 序列长度,长上下文时它是最大瓶颈。
- 追问:和混合精度冲突吗? 不冲突,正交叠加,都是省显存的不同维度。
八、加强记忆
梯度检查点记「少存多算、时间换显存」:反向要用前向的激活,普通做法全存下来(激活显存 O(层数)、长序列爆炸);梯度检查点只存少数检查点、其余丢弃,反向时从最近检查点重新前向算出激活,把激活显存降到约 O(√L),代价是多约一次前向(+~30% 计算),且结果精确不掉点。牢记它只治「激活」这一块,要和混合精度、ZeRO、并行组合成完整的显存优化拳;长上下文训练因激活随序列长度暴涨,它几乎是必开项。