Gradient Checkpointing
了解梯度检查点如何通过在反向传播期间重新计算激活值来减少 GPU 内存占用,并查看 PyTorch 示例、权衡因素和实用训练指南。
梯度检查点是一种节省内存的训练技术,只存储前向传播中的部分中间激活值,并在反向传播期间重新计算其余激活值。它也称为激活检查点,通过增加计算量来降低峰值内存用量。尽管名称中含有“检查点”,这项技术检查点的是激活值,而不是参数梯度或模型文件;当激活张量导致神经网络无法装入可用的 GPU 内存时,这项技术尤其有用。
梯度检查点的工作原理#
在标准前向传播期间,神经网络会计算称为激活值的中间张量。自动微分系统会保留稍后计算梯度所需的激活值,正如 PyTorch 自动微分机制中所述。深层网络、大批次、高分辨率图像和长输入序列都可能导致这些已保存张量占用大量内存。
梯度检查点会将网络划分为多个片段:
- 前向传播会为选定片段保存输入或边界激活值。
- 这些片段内部的其他中间激活值会被丢弃。
- 反向传播期间,每个检查点片段都会再次执行前向传播,以重建缺失的值。
- 重建的激活值会立即用于计算梯度。
PyTorch 激活检查点 API通过 torch.utils.checkpoint 提供这一功能。类似概念还包括 JAX 梯度检查点与重计算和 TensorFlow 磁带检查点。
检查点位置决定了这种权衡。检查点覆盖的区域越多,通常节省的内存越多,但重复执行的运算也越多。与重新计算整个网络相比,在激活值占用内存较多的模块周围选择性地设置检查点,往往能取得更好的平衡。
权衡与相关技术#
梯度检查点通常不会改变模型参数、梯度和优化器状态。它主要针对激活值内存;由于部分前向计算会执行两次,训练速度会变慢。实际效果取决于架构、检查点边界、批次形状和硬件。
它与几种相关技术有所不同:
- 梯度累积会先处理多个微批次,再更新模型权重,从而在不同时加载所有样本的情况下增大有效批次。梯度检查点则会减少每个微批次保留的激活值。
- 混合精度会对选定运算使用较低精度的数据类型。PyTorch 自动混合精度工作流可以减少内存用量并加速兼容的运算,而检查点则有意增加计算量。
- 减小批次大小通过减少同时处理的样本数来降低内存用量。使用检查点则可能让更大的批次或输入分辨率仍然可行。
- 训练检查点会保存权重和优化器状态,以便恢复训练或稍后进行推理。PyTorch 模型检查点指南介绍了这种持久化机制;它与激活值重计算无关。
检查点代码在原始前向执行和重计算前向执行时应保持功能一致。在检查点区域内使用可变状态、设备间传输或不受控的随机性,可能导致错误或不正确的梯度。在解码器式 Transformer 中,检查点训练期间可能还需要禁用键值缓存,因为缓存的推理状态可能与重建前向计算图发生冲突。
PyTorch 示例#
下面的示例展示了如何在一个训练步骤中对占用大量内存的模块使用检查点:
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
torch.manual_seed(0)
block = nn.Sequential(
nn.Linear(1024, 4096),
nn.ReLU(),
nn.Linear(4096, 1024),
)
optimizer = torch.optim.AdamW(block.parameters())
inputs = torch.randn(8, 1024, requires_grad=True)
targets = torch.zeros_like(inputs)
optimizer.zero_grad(set_to_none=True)
outputs = checkpoint(block, inputs, use_reentrant=False)
loss = nn.functional.mse_loss(outputs, targets)
loss.backward()
optimizer.step()只会保留模块的输入和必要的边界信息;其内部激活值会在 loss.backward() 期间重建。实际项目应使用诸如 PyTorch GPU 峰值已分配内存之类的指标,比较启用和未启用检查点的运行情况。
实际应用#
-
高分辨率计算机视觉:医学分割、航空影像检测和工业检测可能需要使用大图像进行训练,而这些图像的特征图占用的内存可能超过模型权重。对选定的主干网络阶段使用检查点,可以保留图像分辨率或允许每批处理更多样本,而不必大幅缩小输入。
-
长序列 Transformer 训练:随着序列长度和层数增加,激活值存储量会迅速增长。重新计算 Transformer 模块可以让更长的上下文或更大的微批次适配同一加速器。NVIDIA 激活值重计算指南展示了对 Transformer 层进行完整和选择性重计算的方法。
实践建议#
当性能分析表明内存主要被激活值而非参数或优化器状态占用时,可以使用梯度检查点。先从较大的重复模块入手,在扩大覆盖范围前,基准测试峰值内存、迭代时间和验证表现。
对于文档齐全的 Ultralytics YOLO 训练配置,优先使用的内存控制手段包括自动混合精度、图像尺寸和物理批次大小。Ultralytics AutoBatch 参考文档介绍了如何根据可用 GPU 内存自动选择批次大小。如果本地硬件仍然不够用,Ultralytics Platform 云端训练可为托管训练运行提供可配置的云 GPU。当自定义 PyTorch 架构需要更精细地管理激活值内存时,梯度检查点可以与这些控制手段配合使用。









