torch.utils.checkpoint是最直接、最低侵入性的显存优化解法,通过重算激活值节省40%~60%显存,代价是训练变慢20%~30%;需严格遵循use_reentrant=false、纯函数封装、避免batchnorm/rnn等副作用模块,并确保输入确定性。

显存撑不住大模型时,torch.utils.checkpoint 是最直接、最低侵入性的解法——它不改模型结构、不依赖多卡、不强制分布式,只靠重算激活值换出 40%~60% 显存,代价是训练变慢约 20%~30%。
什么时候该用 checkpoint() 而不是 checkpoint_sequential()
选 checkpoint() 当你明确知道哪几层“吃内存最多”,且它们能封装成独立函数;选 checkpoint_sequential() 当模型是纯 nn.Sequential 或模块列表,想按层数粗粒度切分(比如每 4 层一组)。
-
checkpoint()更灵活:可嵌套、可带条件逻辑、支持自定义前向逻辑(如跳过 dropout)、适配自定义forward方法 -
checkpoint_sequential()更傻瓜:只认顺序执行的模块,不能处理分支(if)、循环(for)、或跨层依赖(如残差连接需额外传参) - 两者都要求
use_reentrant=False(PyTorch ≥ 1.11 默认),否则在含torch.compile或嵌套 checkpoint 场景下会报RuntimeError: Trying to backward through the graph a second time
checkpoint() 的典型误用和修复方式
常见错误不是语法错,而是语义错:函数内部偷偷用了外部变量、没传全依赖输入、或 RNG 状态没对齐。
快速生成专业的 Python 脚本和应用代码。一键创建完整项目结构,支持CLI、API、爬虫、Bot、Django等多种项目类型,包含完整的项目结构、配置文件、依赖管理、测试、README和文档。
- 错误写法:
def block(x): return self.norm(x) + self.ffn(x)——self是闭包变量,反向重算时拿不到当前self.norm状态 → 改为显式传参:def block(x, norm, ffn): return norm(x) + ffn(x) - 漏传 dropout mask:若函数内有
nn.Dropout,必须设preserve_rng_state=True(默认),否则重算时 dropout 行为不一致 → 不要手动关掉它 - 输入张量未 detach 或 requires_grad=False:
checkpoint要求所有*args是叶子张量且requires_grad=True,否则报Expected all tensors to require gradients→ 检查输入是否来自torch.no_grad()上下文
与 FSDP / DDP 混用时的关键约束
checkpoint 可以和 FSDP 共存,但顺序和位置很关键:checkpoint 必须包裹在 FSDP 包装之后的子模块里,不能包裹整个 FSDP 实例。
- 正确:先
model = FSDP(model),再对model.transformer.layer[5]单独加checkpoint - 错误:对原始
model加checkpoint后再喂给FSDP→ FSDP 初始化阶段会因计算图未就绪而失败 - 和
DDP混用无硬性冲突,但注意 DDP 的find_unused_parameters=True可能与 checkpoint 内部的梯度路径检测冲突 → 建议关掉,改用gradient_checkpointing_enable()(Hugging Face Transformers 封装版) - 性能提醒:checkpoint + FSDP 会放大通信等待时间,因为重算激活期间 GPU 空转,建议配合
backward_prefetch=BackwardPrefetch.BACKWARD_PRE缓解
真正容易被忽略的点
checkpoint 不节省权重显存,只省中间激活;如果你的模型本身参数就超显存(比如 7B 模型放不进 24G 卡),它救不了你——这时候必须上 FSDP、TP 或量化。另外,torch.compile 和 checkpoint 默认不兼容,除非显式启用 use_reentrant=False 并禁用部分优化(如 dynamic=True 会报错)。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










