观文听傑

返回

上一篇用 Automatic Mixed Precision(自动混合精度,AMP)降低了部分张量的位宽,但 Transformer 的长序列激活仍可能比参数更早撑满显存。减小 batch 能继续训练,却会降低设备利用率并改变优化噪声;把张量转成半精度也不能消除随层数和序列长度增长的中间结果。

本篇只解决一个核心问题:Activation Checkpointing(激活检查点)怎样丢弃一部分前向激活,在反向时重算它们,用额外计算换峰值显存。 它不是保存到磁盘的训练 checkpoint,后者用于崩溃恢复。

01 反向传播为什么留着前向激活?#

设一层为

H1=H0W1,H2=GELU(H1),H3=H2W2H_1=H_0W_1,\qquad H_2=\operatorname{GELU}(H_1),\qquad H_3=H_2W_2

其中 H0RB×L×DH_0\in\mathbb{R}^{B\times L\times D}W1RD×4DW_1\in\mathbb{R}^{D\times4D}H1,H2RB×L×4DH_1,H_2\in\mathbb{R}^{B\times L\times4D}W2R4D×DW_2\in\mathbb{R}^{4D\times D}H3RB×L×DH_3\in\mathbb{R}^{B\times L\times D}。反向计算 L/W2=H2(L/H3)\partial\mathcal L/\partial W_2=H_2^\top(\partial\mathcal L/\partial H_3) 需要 H2H_2;GELU 的导数又需要前向输入 H1H_1

Autograd(自动微分)因此保存 backward 所需张量。它们不是参数,也不是最终输出,却会一直活到相应反向算子执行。

flowchart LR
  X[边界输入 H0] --> A[Linear 1]
  A --> B[GELU]
  B --> C[Linear 2]
  C --> Y[边界输出 H3]
  A -.普通前向保存 H1.-> M[(activation memory)]
  B -.普通前向保存 H2.-> M
  Y --> G[backward 到达本段]
  G -->|checkpoint: 从 H0 重跑| A2[重算 Linear 1 + GELU]
  A2 --> R[得到所需 H1/H2]
  R --> D[计算参数与输入梯度]
mermaid

02 Checkpoint 究竟保存和丢弃什么?#

对一段函数 FF 使用 checkpoint 时,第一次前向仍计算 Y=F(X)Y=F(X),但不保留段内所有中间激活;只保留边界输入等重算所需信息。反向到达这段时再次执行 F(X)F(X),重建需要的激活,再照常求梯度。

普通训练: forward(F) -> 保存 a,b,c -> backward(F)
检查点训练: forward(F) -> 只留边界 x -> recompute(F) -> backward(F)
text

被丢掉的是可由边界输入重建的激活,不是参数、参数梯度、optimizer state,也不是当前 batch。若显存主要被 Adam 状态或巨型 embedding 占据,activation checkpointing 的收益就有限。

03 用四层的极小例子手算交换#

假设四层 f1,,f4f_1,\ldots,f_4 的每个边界激活都占 10 MB,忽略参数和临时 workspace。

  • 普通前向保留四层所需激活,粗略为 4×10=404\times10=40 MB;每层前向执行 1 次。
  • f1,f2f_1,f_2 作为一段、f3,f4f_3,f_4 作为一段,只保留两个段输入,边界约 20 MB;反向时两个段各重算 1 次。
  • 若每一层都切成独立 checkpoint,边界本身也要保存,切得越碎不一定继续省;调用和 RNG 管理开销反而增加。

理想化的均匀链可以用约 O(n)O(\sqrt n) 个边界换取额外前向计算,但真实 Transformer 有 attention、MLP、残差分支和不同大小的临时张量,不能只按“层数均分”。应以 profiler 的实际 saved tensor 和峰值为准。

04 当前 PyTorch 的最小正确写法#

PyTorch 2.14 的 torch.utils.checkpoint.checkpoint 要求显式传 use_reentrant,官方推荐 use_reentrant=False。非重入实现会记录前向 autograd graph,并在所需中间量重建完后提前停止重算。

输入 x、参数和返回值的 shape 都不变;变化只发生在 autograd 保存策略和反向执行次数。先在同一模型、同一 batch 上验证 loss 与梯度,再衡量显存和吞吐。

05 为什么边界应包住完整的计算段?#

一个 Transformer block 常包含 LayerNorm、attention、MLP 与残差。只 checkpoint 一个很小的 GELU,保存的边界张量可能和丢掉的激活一样大;把几十层整个包成一段,又会在反向重算过长路径。

常用起点是“每个 block 一段”或“每 2–4 个 block 一段”,随后测量:

切法边界数量重算粒度常见结果
不切0最快,激活显存最高
每个 block易实现,调用开销较多
每 2–4 block常是吞吐与显存折中
整个 stack重算峰值和时延可能过大

应重点覆盖 B×L×4DB\times L\times4D 的 MLP 激活、attention 中随 LL 增长的张量;不要凭模块名猜显存。

06 Dropout 重算为何可能得到另一张图?#

Checkpointed function(被检查点函数)在前向和反向重算时必须等价。Dropout 会消费随机数:若两次 mask 不同,重算的是另一个函数,梯度就不再对应原前向。

默认 preserve_rng_state=True 会保存并恢复 CPU 与一个推断出的设备类型的随机数状态,使重算沿用相同随机结果,但会增加开销。只有当该段确定没有随机算子,或你明确接受不同随机轨迹时,才考虑设为 False

def segment(x, block):                 # block 内可能有 dropout
    return block(x)

y = checkpoint(
    segment, x, block,
    use_reentrant=False,
    preserve_rng_state=True,
)
python

若函数内部把张量移动到运行时新设备,官方文档提醒 RNG 状态可能无法完整预见;应把设备迁移放在 checkpoint 外。

07 副作用和可变全局状态为什么危险?#

反向时函数会再执行一次。因此下面这些副作用可能发生两遍:

  • forward 内追加 Python 列表或递增全局计数器;
  • 更新不受正确控制的 cache;
  • 读取前向后已经改变的配置开关;
  • 让数据相关控制流在两次执行走不同分支。

BatchNorm 的 running statistics、带状态的稀疏路由器和自定义随机 kernel 都应做针对性检查。最安全的 checkpointed segment 是输入相同就产生相同计算图的纯函数式区域。

08 AMP、编译和分布式训练怎样组合?#

Checkpointing 与 AMP 解决不同维度:前者减少保存的激活,后者改变部分算子的 dtype。重算必须处于与原前向相容的 autocast 上下文;不要让重算偷偷用另一种精度。框架封装模型时,应做固定 batch 的 FP32、AMP、AMP+checkpoint 三路对照。

DDP(Distributed Data Parallel,分布式数据并行)下,每个 rank 都在本地重算;通信量通常不因此减少。torch.compile、FSDP 与 selective checkpointing 会改变图捕获或保存策略,组合后要重新 profile,不能把各自节省比例直接相乘。

09 如何测到真正的峰值显存与代价?#

CUDA 是异步的,单看某一行之后的 memory_allocated() 容易误判。至少预热若干步,再同步并记录完整训练 step。

torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)

start.record()
loss = train_step(batch)
end.record()
torch.cuda.synchronize()

peak = torch.cuda.max_memory_allocated() / 2**30
step_ms = start.elapsed_time(end)
print({"loss": float(loss), "peak_GiB": peak, "step_ms": step_ms})
python

比较时固定模型、序列长度、micro-batch、dtype、编译设置和随机种子。报告有效 tokens/s、峰值 allocated/reserved 显存、重算算子时间与最终验证指标;只说“省了 40%”没有可迁移意义。

10 常见错误与最短调试路径#

症状常见原因最短检查
显存几乎没降峰值来自参数/optimizer 或边界太碎拆分 memory snapshot,统计 saved tensors
loss 相同但梯度不同Dropout RNG 未保持、全局状态改变关随机层后复现,再检查 RNG 与副作用
反向突然很慢段过大或重复重算昂贵算子profiler 标记 recompute 区域
checkpoint 后报 shape/dtype 不同两次执行走了不同分支debug=True,固定输入和配置
自定义 backward 失效重入实现限制或图契约不兼容显式改 use_reentrant=False 做最小例
仍然 OOMattention 临时量或通信 buffer 才是峰值逐阶段测峰值,而非只看 forward 末尾

非重入实现的默认 determinism check 会比较重算张量的 shape、dtype 与 device,但它不是数值相等证明。关键实验仍应比较参数梯度的有限性、方向与短程收敛。

11 它和哪些“检查点”不是一回事?#

方法保存位置解决的问题主要代价
Activation checkpointing运行时边界张量单步激活显存额外重算
Training checkpoint磁盘/对象存储故障恢复、续训I/O 与存储
CPU/NVMe offload主存/磁盘GPU 常驻显存传输时延
Gradient accumulation参数 .grad有效 batch 大小更多 micro-step

Activation checkpointing 不会替你保存 optimizer、scaler、token 时钟或数据游标;机器中断后仍需要上一篇所示的训练 checkpoint。

12 失败边界#

当模型是计算密集型且显存只差一点时,重算通常值得;当训练已受算力限制、模块极小、数据管线空转,额外 forward 可能让吞吐下降得更多。包含不可重放 I/O、跨设备随机状态或不可重复副作用的区域不适合直接 checkpoint。

它也不会降低推理 KV Cache,因为推理没有同样的反向图;推理显存应从 cache 长度、量化、分页管理或并行策略入手。

13 今天真正需要记住什么?#

  1. 反向需要前向激活;checkpoint 只保留边界,反向到达时重算段内中间量。
  2. 省下的是激活显存,代价是额外计算;切分边界必须用 profiler 验证。
  3. PyTorch 2.14 应显式使用 use_reentrant=False;随机层默认保留 RNG 状态。
  4. 前向与重算必须等价,副作用、设备迁移和动态分支是高风险区。

14 思考题与小练习#

  1. 对 12 个等成本 block,分别画出每层一段、每 3 层一段时保存的边界和反向重算顺序;估算每种方案的额外 forward 次数。
  2. 给含 Dropout 的两层 MLP 写一个梯度对照测试:普通、preserve_rng_state=TrueFalse 三种配置比较同一参数的梯度最大绝对误差。
  3. 写一个 benchmark 同时记录峰值显存、step time 与 tokens/s,并解释为什么必须预热和 torch.cuda.synchronize()

相关工作#

  1. Chen et al., Training Deep Nets with Sublinear Memory Cost,系统展示用重算把深网训练的激活内存降到次线性规模。
  2. Griewank & Walther, Algorithm 799: Revolve,研究受限 checkpoint 数量下的最优反向重算调度。
  3. Jain et al., Checkmate: Breaking the Memory Wall with Optimal Tensor Rematerialization,把计算图上的重算边界选择建模为优化问题。
  4. Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models,给出 Transformer 激活内存组成与选择性策略的工程分析。

15 下一篇预告#

显存允许更长的计算图后,有效 batch 仍可能无法一次塞进设备。下一篇将把一个大 batch 拆成多个 micro-batch,解释梯度累积何时与一次性训练等价,以及 DDP 中如何只在最后一次反向同步。

激活占满显存时该丢掉什么?Activation Checkpointing 的重算、边界与随机数状态
https://zwjcode.cn/blog/activation-checkpointing-recompute-memory-boundary
作者
发布于 2026年9月14日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。