← 返回题目列表

什么是梯度检查点(激活重计算)?它如何用时间换显存?

高频 中等 第 9 / 25 题 更新于 2026/07/28
梯度检查点激活重计算显存优化反向传播

简化版

梯度检查点(gradient checkpointing,又叫激活重计算) 是一种「用计算换显存」的技巧。反向传播需要前向时的激活值来算梯度,普通做法是把每一层的激活都存下来,显存开销随层数和序列长度线性增长、非常大。梯度检查点的做法是:前向时只保存少数几个”检查点”层的激活,其余丢弃;反向传播到某段时,从最近的检查点重新前向计算出需要的激活。这样激活显存从 O(层数) 降到约 O(√层数),代价是多做一次前向(约增加 30% 计算量)。它是训练长序列/大模型时省激活显存的常用手段。

详细版

为什么激活这么占显存

反向传播算梯度用到链式法则,需要前向时每一层的中间激活。训练时的显存主要由四块组成:参数、梯度、优化器状态、激活。其中激活显存 ∝ batch × 序列长度 × 隐藏维 × 层数,在长序列/大 batch 时会非常庞大,常常是显存瓶颈。

梯度检查点的做法

普通前向:保存每一层的激活 → 反向直接用      → 激活显存 O(L)
检查点:  只保存 √L 个检查点的激活 → 反向时从检查点重算中间激活 → 激活显存 O(√L)
  • 把网络分成若干段,只在段边界保存激活(检查点)。
  • 反向传播进入某段时,以该段起点的检查点为输入,重新前向一遍算出段内各层激活,再算梯度。
  • 用完即丢,不长期占显存。

代价与收益

指标普通梯度检查点
激活显存O(L)约 O(√L)
计算量1 次前向 + 1 次反向约 2 次前向 + 1 次反向(多 ~30%)

适用场景

  • 长上下文训练、大 batch、单卡显存吃紧时。
  • 和混合精度、ZeRO、并行等叠加使用。

完整版教学

一、先看清训练显存都花在哪

训练一个模型,显存主要被四部分吃掉:

  1. 参数:模型权重。
  2. 梯度:与参数等大。
  3. 优化器状态:Adam 的动量 m、方差 v、FP32 副本等(往往最大头之一)。
  4. 激活(activations):前向过程中每一层产生的中间结果,反向要用它们算梯度。

前三块由参数量决定,相对固定;而激活显存随 batch、序列长度、层数线性增长,在长序列训练里会爆炸式增大,常常成为压垮单卡的最后一根稻草。梯度检查点专治这一块。

二、为什么反向传播非要激活不可

反向传播用链式法则逐层往回算梯度。算某一层参数的梯度时,公式里会用到该层前向时的输入/中间激活。举个直觉例子:一个乘法 y = w·x∂y/∂w = x——要算对 w 的梯度,就必须知道当时的输入 x。所以普通实现会在前向时把所有层的激活缓存下来,供反向使用。层越多、序列越长,这个缓存越大。

三、核心思想:不存就重算

梯度检查点的洞察是:激活是可以”重新算出来”的——只要保留了某个位置的输入,就能通过前向计算再得到后面的激活。既然如此,何必全都存着占显存?不如丢掉大部分,需要时重算

具体做法(以按段划分为例):

  1. 把 L 层网络切成若干段,只在段的边界保存激活(这些就是”检查点”)。
  2. 前向时,段内的中间激活算完就丢弃,不缓存。
  3. 反向传播到某段时,取该段起点的检查点作为输入,重新前向计算一遍这一段,临时得到段内激活,用来算梯度,算完再丢。

这样常驻的激活显存只剩下「检查点」的量,大幅下降。

四、为什么是 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 spikeoverflow/NaN
语料 -> 清洗去重 -> tokenization -> 分布式训练 -> checkpoint -> 评测
  |        |             |                |             |
 质量     覆盖率         吞吐              可恢复         能力验证

所以回答「什么是梯度检查点(激活重计算)?它如何用时间换显存?」时,不能只说某个技巧“省显存”或“加速”。要说明它省的是哪一笔账、牺牲了什么、线上训练日志里应该观察哪个信号。

七、常见误区与追问

  • 误区:梯度检查点能省参数/优化器显存。 不能,它只省激活显存;参数/梯度/优化器状态要靠 ZeRO、混合精度等。
  • 误区:它会降低模型精度。 不会,重算得到的激活和原来完全一致,是精确的,只是多花计算。
  • 追问:为什么反向需要激活? 链式法则里参数梯度依赖前向的输入/中间值(如 ∂(wx)/∂w = x)。
  • 追问:省了多少、代价多少? 激活显存 O(L)→约 O(√L),计算量约增 30%(多一次前向)。
  • 追问:为什么长序列训练一定要它? 激活显存 ∝ 序列长度,长上下文时它是最大瓶颈。
  • 追问:和混合精度冲突吗? 不冲突,正交叠加,都是省显存的不同维度。

八、加强记忆

梯度检查点记「少存多算、时间换显存」:反向要用前向的激活,普通做法全存下来(激活显存 O(层数)、长序列爆炸);梯度检查点只存少数检查点、其余丢弃,反向时从最近检查点重新前向算出激活,把激活显存降到约 O(√L),代价是多约一次前向(+~30% 计算),且结果精确不掉点。牢记它只治「激活」这一块,要和混合精度、ZeRO、并行组合成完整的显存优化拳;长上下文训练因激活随序列长度暴涨,它几乎是必开项。