什么是梯度检查点
·
梯度检查点(Gradient Checkpointing) 是一种在训练大模型时常用的 显存优化技术,其目标是在不牺牲模型性能的前提下,节省 GPU 显存,以便能训练更大的模型或使用更大的 batch size。
🧠 原理概述
在标准的训练过程中(使用自动微分),前向传播(forward)时每层的中间激活值都会保留,以便在反向传播(backward)时计算梯度。
问题是:
这些中间值占用了大量显存,尤其是在深层模型中。
✅ 梯度检查点怎么做?
梯度检查点的核心思想是:
在前向传播时只保留部分关键节点(称为“检查点”)的中间激活值,其余的在反向传播时再重新计算。
也就是说:
-
少存一点中间结果(省显存)。
-
回头再算一次前向(多花一点算力)。
是一种 以计算换显存 的策略。
🧮 举个例子
假设你有一个模型分成了 10 层:
-
正常训练时:forward 会保存 10 层的中间激活,占用大量显存。
-
用梯度检查点:比如只保存第 0 层、第 5 层、第 10 层的结果。
-
backward 时,先从第 5 层重算 6~10 层,再计算这些层的梯度。
-
这样最多只需要保存一小部分层的中间结果,节省大量显存。
-
📉 优缺点
| 项目 | 优点 | 缺点 |
|---|---|---|
| 显存使用 | 显著减少 | — |
| 计算量 | 增加(重复前向) | 更慢一些 |
| 适用场景 | 大模型、内存受限时训练 | 对推理没用 |
🔧 在 PyTorch 中的用法示例
import torch
from torch.utils.checkpoint import checkpoint
def custom_forward(*inputs):
return model_layer(*inputs)
# 使用 checkpoint 包裹某一层
output = checkpoint(custom_forward, input_tensor)
更多推荐



所有评论(0)