梯度检查点与激活重算 本节摘要:反向传播要保留每一层的中间激活。在 70B 参数、128K 上下文的规模下,这是每张卡 3TB 的激活量。梯度检查点(gradient checkpointing, aka 激活重算 activation recomputation)用「算力换显存」——不存中间激活,反向时重算来恢复。本节讲透它为何必要、朴素全检查点的 33% 额外 FLOPs 代价、Korthikanti(2022)选择性检查点如何做到 5 倍省显存且仅 5% 额外开销,以及 CPU 卸载(offload)与混合策略。读完本节,你会理解长上下文训练里「注意力 softmax 的 O(L²) 激活」才是真正的大头,以及为何 2024+ 前沿训练几乎都用选择性重算。
本节摘要:反向传播要保留每一层的中间激活。在 70B 参数、128K 上下文的规模下,这是每张卡 3TB 的激活量。梯度检查点(gradient checkpointing, aka 激活重算 activation recomputation)用「算力换显存」——不存中间激活,反向时重算来恢复。本节讲透它为何必要、朴素全检查点的 33% 额外 FLOPs 代价、Korthikanti(2022)选择性检查点如何做到 5 倍省显存且仅 5% 额外开销,以及 CPU 卸载(offload)与混合策略。读完本节,你会理解长上下文训练里「注意力 softmax 的 O(L²) 激活」才是真正的大头,以及为何 2024+ 前沿训练几乎都用选择性重算。
对应原课程:Phase 10 · Lesson 24 ·
34-gradient-checkpointing(原英文phases/10-llms-from-scratch/34-gradient-checkpointing/docs/en.md)。
阅读完本节,你应当能够:
训练一个 Transformer,反向传播要为每一层保留每个可微操作的输入:注意力输入、Q/K/V 投影、softmax 输出、FFN 输入、归一化输出、残差流。对隐藏维度 d、序列长 L、batch B 的一层,这大概是 12*B*L*d 个 float。
d=8192, L=8192, B=1 时,一层就是 BF16 下 800MB;64 层模型是 51GB 激活——这还没乘 microbatch、还没加注意力的 L² 中间量、还没算张量并行的副本。
两面夹击:权重加优化器状态或许能塞进 80GB,但激活把你挤过线。梯度检查点是标准解法——丢掉大部分激活,反向时重做前向把它们算回来。代价是额外 FLOPs,收益是显存按「检查点段数/总层数」的比例下降。朴素做约 33% 额外前向 FLOPs;Korthikanti 的「聪明选择」能做到 5 倍省显存且开销低于 5%。
output = layer(input)。反向要 grad_input 和 grad_params,为此需要:
input(线性层算 grad_params = input.T @ grad_output)前向自动把这些存进 autograd 图。
把网络分成 N 段。前向只存每段的输入;反向需要中间量时,重跑该段前向把它们物化,再求导。
例:32 层 Transformer 分成 32 段(每段 1 层):
这是 Chen 等 2016 的原始配方:每 sqrt(L) 层一个检查点以平衡显存与算力。L=64 时即 8 个检查点。
并非所有激活代价相同:
B*L*L*heads,随序列长二次方增长。B*L*4d,线性增长。长序列下 softmax 占主导。选择性检查点保留便宜的激活(线性投影、残差),只重算昂贵的(注意力)。你付出极少 FLOPs 重算,却省下 O(L²) 显存。Megatron-Core 把它实现为「选择性激活重算」,2024+ 几乎所有前沿训练都在用。
朴素每 k 层一检查点(L 层共)的每步 FLOPs:
flops_fwd_normal = L * f_layer flops_bwd_normal = 2 * L * f_layer flops_total_normal = 3 * L * f_layer flops_recompute = L * f_layer # 每段多一次前向 flops_total_ckpt = 4 * L * f_layer overhead = 4/3 - 1 = 33%
选择性检查点只重算注意力内核:
flops_recompute_selective = L * f_attention ≈ L * f_layer * 0.15 overhead_selective = (3 + 0.15)/3 - 1 = 5%
💡 5% 开销换 5 倍显存:选择性检查点是长上下文训练的「免费午餐」——只重算占比小但显存占比大(softmax)的部分。
重算的替代:前向与反向之间把激活送到 CPU 内存。需要 PCIe 带宽;当空闲带宽超过重算代价时有益。混合策略常见:某些层检查点、某些层卸载。FSDP2 把卸载作为一等选项,在「GPU 显存瓶颈但 CPU-GPU 传输有余量」时大放异彩。
| 策略 | 显存节省 | 额外 FLOPs | 适用 |
|---|---|---|---|
| 不检查点 | 无 | 0 | 短序列、显存充足 |
| 朴素全检查点 | 大(按段数比例) | ~33% | 显存极紧、算力富余 |
| 选择性检查点 | 大(省 O(L²)) | ~5% | 长上下文(主流) |
| CPU 卸载 | 大 | 0 FLOPs,耗 PCIe 带宽 | 显存紧、带宽有余 |
| 混合 | 可调 | 可调 | 生产精细调优 |
PyTorch 的 torch.utils.checkpoint、Megatron-Core、FSDP2 都提供这些开关。
本节产出 outputs/skill-activation-recomputation.md——显存不够时的决策树:先选择性检查点(5% 开销)→ 仍不够加朴素检查点 → 仍不够加 CPU 卸载 → 仍不够上 ZeRO-3/FSDP 分片。附带「先省 O(L²) 再省 O(L)」的优先级心法。
torch.utils.checkpoint.checkpoint 包裹每层,测量显存下降与训练速度变化,对比 33% 理论开销。12*B*L*d,长上下文 + 大模型下爆炸(70B/128K → 3TB/卡)。sqrt(L) 层一点。本节是「从零构建 LLM」章的收尾。回到「教程总纲.md」可进入下一篇(大模型工程)的 LLM 工程化内容。