什么是梯度检查点(激活重计算)?它如何用时间换显存?
简化版
梯度检查点(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:分片参数/梯度/优化器状态,省前三块的冗余。
- 张量/流水线并行:把模型拆到多卡。
组合拳里,梯度检查点是「针对激活、几乎零副作用(只多点计算)」的常用项,尤其在长上下文训练时几乎必开——因为长序列下激活显存增长最凶。
六、面试拆解算例
预训练题最好落到预算账和稳定性账。假设训练一个 7B 模型,目标 token 数是 1T,按常见粗估 6 × 参数量 × token 数,训练计算量约为 6 × 7e9 × 1e12 = 4.2e22 FLOPs。如果有效集群算力是 1e18 FLOPs/s,理想情况下也要约 42,000 秒;现实还要扣通信、数据加载、checkpoint 和故障恢复的损耗。
training_flops ≈ 6 * N * D
N = 7e9 parameters
D = 1e12 tokens
training_flops ≈ 4.2e22
| 账本 | 关键变量 | 常见瓶颈 | 排查信号 |
|---|---|---|---|
| 数据账 | token 数、重复率、质量分 | 脏数据和污染 | eval 异常偏高 |
| 算力账 | GPU 数、利用率、通信 | MFU 低 | step time 抖动 |
| 显存账 | batch、序列长、优化器状态 | OOM | 激活占用过高 |
| 稳定性账 | 学习率、精度、梯度 | loss spike | overflow/NaN |
语料 -> 清洗去重 -> tokenization -> 分布式训练 -> checkpoint -> 评测
| | | | |
质量 覆盖率 吞吐 可恢复 能力验证
所以回答「什么是梯度检查点(激活重计算)?它如何用时间换显存?」时,不能只说某个技巧“省显存”或“加速”。要说明它省的是哪一笔账、牺牲了什么、线上训练日志里应该观察哪个信号。
七、常见误区与追问
- 误区:梯度检查点能省参数/优化器显存。 不能,它只省激活显存;参数/梯度/优化器状态要靠 ZeRO、混合精度等。
- 误区:它会降低模型精度。 不会,重算得到的激活和原来完全一致,是精确的,只是多花计算。
- 追问:为什么反向需要激活? 链式法则里参数梯度依赖前向的输入/中间值(如 ∂(wx)/∂w = x)。
- 追问:省了多少、代价多少? 激活显存 O(L)→约 O(√L),计算量约增 30%(多一次前向)。
- 追问:为什么长序列训练一定要它? 激活显存 ∝ 序列长度,长上下文时它是最大瓶颈。
- 追问:和混合精度冲突吗? 不冲突,正交叠加,都是省显存的不同维度。
八、加强记忆
梯度检查点记「少存多算、时间换显存」:反向要用前向的激活,普通做法全存下来(激活显存 O(层数)、长序列爆炸);梯度检查点只存少数检查点、其余丢弃,反向时从最近检查点重新前向算出激活,把激活显存降到约 O(√L),代价是多约一次前向(+~30% 计算),且结果精确不掉点。牢记它只治「激活」这一块,要和混合精度、ZeRO、并行组合成完整的显存优化拳;长上下文训练因激活随序列长度暴涨,它几乎是必开项。