梯度检查点(激活重计算)
Gradient Checkpointing (Activation Recomputation)进阶前向时少存中间激活,反向时再重算一遍,用计算换显存。
训练时,前向传播产生的中间结果(激活)默认要一直存到反向传播用完,这是显存的大头之一。梯度检查点只在少数检查点位置保存激活,其余的丢掉,反向传播需要时从最近的检查点重新前向计算补回来。2016 年陈天奇等人的论文《Training Deep Nets with Sublinear Memory Cost》系统提出了这个做法:n 层网络只需约 O(√n) 的激活显存,代价是每个小批次多做大约一次前向计算;论文里 1000 层残差网络的显存从 48GB 降到 7GB,运行时间增加约 30%。PyTorch 的 torch.utils.checkpoint 就是它的实现,训练大模型和 VLA 时常与混合精度、梯度累积一起开。
例子微调 Transformer 策略时把每个 Transformer 层包进 torch.utils.checkpoint,显存占用明显下降,每步训练会慢一些。
- 也叫
- 激活检查点、Activation Checkpointing
- 相关
- 反向传播、GPU 显存、梯度累积、混合精度训练、全分片数据并行、DeepSpeed
- 来源
- Training Deep Nets with Sublinear Memory Cost (arXiv 1604.06174)
torch.utils.checkpoint (PyTorch 文档)