← 返回题目列表

Activation Checkpointing 如何节省 Transformer 训练显存?

高频 困难 第 17 / 25 题 更新于 2026/09/17
Activation CheckpointingTransformer显存优化重计算

简化版

Activation Checkpointing 不保存所有层的前向中间激活,只保留若干检查点;反向到某段时,从最近检查点重新执行前向来恢复需要的张量,因此用额外计算换显存。分段越大省得越多、重算越多。实现要正确处理随机数、Dropout、副作用和混合精度,并按峰值显存、tokens/s 和数值一致性选择粒度。

详细版

没有 checkpointing 时 L 层大致保存 O(L) 份激活;均匀分成约 √L 段,可把保存量降到近 O(√L) 的数量级,同时每段前向在反向中重算。现代框架还支持按 Transformer block、Attention 或 FFN 选择性 checkpoint。

forward: save x0 ---- save xk ---- save x2k
backward: recompute segment 2 -> grad
          recompute segment 1 -> grad

重算必须产生与原前向兼容的值,Dropout RNG 状态需要保存/恢复;有文件写入、计数器等副作用的函数不能直接重放。使用非 reentrant/reentrant 接口时要遵循框架限制。评测完整 step time、峰值显存、可增加的 batch/sequence 和最终收敛,而不只看理论节省。

完整版教学

一、为什么激活可以不全部保存

反向需要前向中间值,但这些值可由段入口与参数重新计算。Activation Checkpointing 选择少量边界保留,释放段内激活;反向时再跑一遍该段前向。

它不压缩参数,也不减少数学模型容量。节省的是“为反向暂存什么”,代价是重复计算与可能的通信。

记忆钩子:Checkpointing 像只记章节起点,复习某页时从章节开头重读;少做笔记,多花重读时间。

二、一个三段执行例子

假设 12 个 block 分成 3 段,每段 4 层。首次前向仅保存段边界 x0、x4、x8、x12,段内 x1~x3 等释放。

反向先从 x8 重算第 9~12 层,得到段内激活后求梯度;再从 x4 重算下一段。每个被 checkpoint 的段通常多做一次前向。

x0 -> [1..4] -> x4 -> [5..8] -> x8 -> [9..12] -> x12
       discard           discard            discard

三、显存与计算的数量级权衡

简化地说,不做重算保存 L 层;均匀分段若段长 k,需要保存约 L/k 个边界和重算时 k 个段内激活,总峰值约 L/k+k,在 k≈√L 附近最小。

L=64 时取 k=8,粗略项为 8+8=16,而不是 64。实际 Transformer 各层张量大小和生命周期不同,不能把公式当精确字节。

重算增加的 FLOPs 主要是前向部分,训练墙钟增幅通常小于 100%,具体由 kernel 和通信占比决定。

四、按什么粒度切分

整 block checkpoint 简单稳健;只 checkpoint Attention 或 FFN 可针对大激活;更细粒度能控制重算,但增加图管理开销和复杂性。

粒度显存节省重算/复杂度
每若干 block易实现
每个 block重算更多
选择性算子定向需了解保存张量
逐算子可很高调度开销与维护高

异构层或 MoE 不宜机械等层分段,应按实际激活字节平衡。

五、Dropout 与随机数一致性

首次前向和重算若使用不同 Dropout mask,反向对应的就不是同一函数值,梯度可能错误。框架通常可保存和恢复 RNG 状态,但会有额外开销。

模型并行下不同设备的 RNG 流更复杂,要确认 checkpoint API 是否覆盖所有设备。为了速度关闭 RNG 保留前,必须证明段内没有随机操作。

数值一致不一定逐 bit 相同,至少要在容差内比较 loss 和梯度,并跑收敛回归。

六、重算函数必须近似纯函数

函数若在前向写文件、更新全局计数、消费队列或改变缓存,反向重算会重复副作用。模型 block 通常应只依赖输入、参数与可恢复随机状态。

BatchNorm 运行统计、某些动态路由日志和自定义 kernel 状态需要特别检查。训练监控计数应放在 checkpoint 段外或保证幂等。

输入结构、requires_grad 和 detach 行为也受框架 reentrant 模式限制,应按当前版本文档测试。

七、与并行和通信的交互

张量并行重算某层可能重复 all-reduce/all-gather,使代价高于纯计算。Pipeline 并行的多个 micro-batch 还会改变激活峰值与重算调度。

有些系统保存通信后的结果以避免重通信,只重算本地算子;这是“选择性 checkpoint”的一部分。最优方案需通过 profiler 看计算与通信占比。

ZeRO 切模型状态,checkpointing 切激活,二者互补;CPU offload 则用传输换显存。

八、如何做工程选型

先测无 checkpoint 基线的峰值与 tokens/s,再逐级开启每两层、每层或选择性方案。记录能否提升 micro-batch/sequence,以及提升后整体吞吐。

例如 checkpoint 使单步慢 25%,却允许 batch 从 2 增到 4,tokens/s 可能反而提高;也可能计算瓶颈下净吞吐下降。必须比较最终每 GPU token 速度。

同时检查 loss 曲线、梯度 norm、固定 batch 输出和断点恢复,避免“显存省了但训练悄悄错了”。

九、常见误区与追问

  • 误区:Checkpointing 会减少模型参数。 它只减少保存的激活。
  • 误区:所有中间值都不保存。 仍需保存段边界与非重算状态。
  • 误区:重算只增加少量代码,不影响通信。 并行层可能重复 collective。
  • 误区:Dropout 重算无需保持随机状态。 不同 mask 会破坏梯度对应关系。
  • 误区:粒度越细一定越好。 图开销、重算和维护成本会增加。
  • 追问:理论最优段长是多少? 均匀同尺寸简化模型约 √L,但真实应按字节和 profiler。
  • 追问:训练会慢多少? 取决于 checkpoint 覆盖、通信与 kernel,必须实测 step time。
  • 追问:与梯度累积有何区别? 累积减 micro-batch 峰值并保持全局 batch,checkpoint 重算单个 micro-batch 的层内激活。

十、加强记忆

Activation Checkpointing 记住“存边界、丢中间、反向重算”:它用额外前向与可能的通信换激活显存,均匀分段有 L/k+k 的粗略权衡,但真实粒度由逐层字节决定。Dropout RNG 与函数副作用是正确性红线;最后同时看峰值、step time、可扩 batch/sequence、tokens/s 和收敛,不能只报告省了多少 GB。