直接回答:梯度检查点(也叫激活重算,activation recomputation)用计算换显存。标准反向传播需要保存每一层的中间激活,显存随网络深度线性增长,长序列大模型下激活往往比参数本身还占显存。检查点技术前向时只保存少数边界位置的激活,反向传播到某一段时,从最近的检查点临时重新前向算出该段所需的中间激活再求梯度。理论上按最优分段(大约每 √n 层设一个检查点)可把激活显存从 O(n) 降到 O(√n),典型实现的代价是约 20%–40% 的额外重算时间;省下的显存可以换取更大的 batch 或更长的上下文。

实践要点:PyTorch 用 torch.utils.checkpoint 包裹要重算的子模块即可,新版推荐非重入(non-reentrant)实现,与 FSDP 等分布式方案兼容性更好;进阶做法是选择性检查点(selective checkpointing)——只重算显存大但计算廉价的 element-wise 算子,保留昂贵 matmul 的结果,性价比明显更高,主流训练框架已内置。一个常见坑:dropout 等随机算子在重算时必须复现与前向完全相同的随机数状态,否则重算出的激活与前向不一致、梯度静默出错,框架通过保存 RNG 状态解决,自定义算子需要自行处理。

from torch.utils.checkpoint import checkpoint

# 只对 Transformer block 做检查点,embedding 层正常保存
x = checkpoint(transformer_block, x, use_reentrant=False)

追问方向:选择性检查点如何挑选保留与重算的算子?它与激活量化、CPU offload 如何组合使用?对端到端训练吞吐的实际影响如何测量?(约 560 字)