激活占满显存时该丢掉什么?Activation Checkpointing 的重算、边界与随机数状态
从反向传播为何需要中间激活出发,手算显存—计算交换,拆解 PyTorch 2.14 非重入 checkpoint 的重算数据流,并给出边界选择、随机层、调试与性能验证方法。
上一篇用 Automatic Mixed Precision(自动混合精度,AMP)降低了部分张量的位宽,但 Transformer 的长序列激活仍可能比参数更早撑满显存。减小 batch 能继续训练,却会降低设备利用率并改变优化噪声;把张量转成半精度也不能消除随层数和序列长度增长的中间结果。
本篇只解决一个核心问题:Activation Checkpointing(激活检查点)怎样丢弃一部分前向激活,在反向时重算它们,用额外计算换峰值显存。 它不是保存到磁盘的训练 checkpoint,后者用于崩溃恢复。
01 反向传播为什么留着前向激活?#
设一层为
其中 ,,,,。反向计算 需要 ;GELU 的导数又需要前向输入 。
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[计算参数与输入梯度]mermaid02 Checkpoint 究竟保存和丢弃什么?#
对一段函数 使用 checkpoint 时,第一次前向仍计算 ,但不保留段内所有中间激活;只保留边界输入等重算所需信息。反向到达这段时再次执行 ,重建需要的激活,再照常求梯度。
普通训练: forward(F) -> 保存 a,b,c -> backward(F)
检查点训练: forward(F) -> 只留边界 x -> recompute(F) -> backward(F)text被丢掉的是可由边界输入重建的激活,不是参数、参数梯度、optimizer state,也不是当前 batch。若显存主要被 Adam 状态或巨型 embedding 占据,activation checkpointing 的收益就有限。
03 用四层的极小例子手算交换#
假设四层 的每个边界激活都占 10 MB,忽略参数和临时 workspace。
- 普通前向保留四层所需激活,粗略为 MB;每层前向执行 1 次。
- 把 作为一段、 作为一段,只保留两个段输入,边界约 20 MB;反向时两个段各重算 1 次。
- 若每一层都切成独立 checkpoint,边界本身也要保存,切得越碎不一定继续省;调用和 RNG 管理开销反而增加。
理想化的均匀链可以用约 个边界换取额外前向计算,但真实 Transformer 有 attention、MLP、残差分支和不同大小的临时张量,不能只按“层数均分”。应以 profiler 的实际 saved tensor 和峰值为准。
04 当前 PyTorch 的最小正确写法#
PyTorch 2.14 的 torch.utils.checkpoint.checkpoint 要求显式传 use_reentrant,官方推荐 use_reentrant=False。非重入实现会记录前向 autograd graph,并在所需中间量重建完后提前停止重算。
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
class Block(nn.Module):
def __init__(self, d_model=512):
super().__init__()
self.norm = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, 4 * d_model),
nn.GELU(),
nn.Linear(4 * d_model, d_model),
)
def forward(self, x): # x: [B,L,D]
return x + self.ffn(self.norm(x))
class Stack(nn.Module):
def __init__(self, depth=12, d_model=512):
super().__init__()
self.blocks = nn.ModuleList([Block(d_model) for _ in range(depth)])
def forward(self, x): # [B,L,D] -> [B,L,D]
for block in self.blocks:
x = checkpoint(
block,
x,
use_reentrant=False,
preserve_rng_state=True,
)
return xpython输入 x、参数和返回值的 shape 都不变;变化只发生在 autograd 保存策略和反向执行次数。先在同一模型、同一 batch 上验证 loss 与梯度,再衡量显存和吞吐。
05 为什么边界应包住完整的计算段?#
一个 Transformer block 常包含 LayerNorm、attention、MLP 与残差。只 checkpoint 一个很小的 GELU,保存的边界张量可能和丢掉的激活一样大;把几十层整个包成一段,又会在反向重算过长路径。
常用起点是“每个 block 一段”或“每 2–4 个 block 一段”,随后测量:
| 切法 | 边界数量 | 重算粒度 | 常见结果 |
|---|---|---|---|
| 不切 | 0 | 无 | 最快,激活显存最高 |
| 每个 block | 多 | 细 | 易实现,调用开销较多 |
| 每 2–4 block | 中 | 中 | 常是吞吐与显存折中 |
| 整个 stack | 少 | 粗 | 重算峰值和时延可能过大 |
应重点覆盖 的 MLP 激活、attention 中随 增长的张量;不要凭模块名猜显存。
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 做最小例 |
| 仍然 OOM | attention 临时量或通信 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 今天真正需要记住什么?#
- 反向需要前向激活;checkpoint 只保留边界,反向到达时重算段内中间量。
- 省下的是激活显存,代价是额外计算;切分边界必须用 profiler 验证。
- PyTorch 2.14 应显式使用
use_reentrant=False;随机层默认保留 RNG 状态。 - 前向与重算必须等价,副作用、设备迁移和动态分支是高风险区。
14 思考题与小练习#
- 对 12 个等成本 block,分别画出每层一段、每 3 层一段时保存的边界和反向重算顺序;估算每种方案的额外 forward 次数。
- 给含 Dropout 的两层 MLP 写一个梯度对照测试:普通、
preserve_rng_state=True、False三种配置比较同一参数的梯度最大绝对误差。 - 写一个 benchmark 同时记录峰值显存、step time 与 tokens/s,并解释为什么必须预热和
torch.cuda.synchronize()。
相关工作#
- Chen et al., Training Deep Nets with Sublinear Memory Cost ↗,系统展示用重算把深网训练的激活内存降到次线性规模。
- Griewank & Walther, Algorithm 799: Revolve ↗,研究受限 checkpoint 数量下的最优反向重算调度。
- Jain et al., Checkmate: Breaking the Memory Wall with Optimal Tensor Rematerialization ↗,把计算图上的重算边界选择建模为优化问题。
- Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models ↗,给出 Transformer 激活内存组成与选择性策略的工程分析。
15 下一篇预告#
显存允许更长的计算图后,有效 batch 仍可能无法一次塞进设备。下一篇将把一个大 batch 拆成多个 micro-batch,解释梯度累积何时与一次性训练等价,以及 DDP 中如何只在最后一次反向同步。