torch.utils.checkpoint.checkpoint 可安全省显存,但必须用 use_reentrant=false、纯函数封装、避免 batchnorm/rnn 等副作用模块,并确保输入确定性;否则易报错或无效。

能省,但必须用对函数、参数和模式——否则显存没降,反而报错或训练失效。
torch.utils.checkpoint.checkpoint 函数怎么调才安全
别直接包装整个 nn.Module 实例,它会隐式捕获 self.training 等状态,导致重算时行为不一致。正确做法是把前向逻辑抽成纯函数,只接收 Tensor 参数:
- 函数内部不能调用
nn.Dropout(训练态随机不可复现),改用nn.functional.dropout(x, training=False)或移出检查点范围 - 避免闭包变量,比如不要在函数里读
self.hidden_size;所有非Tensor参数必须显式传入,如checkpoint(custom_forward, x, layer, dropout_p=0.1) - 确保函数可重入:同一输入、同一设备、同一 dtype 下,多次调用输出完全一致
use_reentrant=False 是强制要求,不是可选项
PyTorch 2.5+ 已弃用 use_reentrant=True 模式,强行使用会触发警告甚至报错。Non-reentrant 模式通过张量钩子保存输入,不再依赖计算图重放,更稳定且显存节省更高(实测多降 10–15%):
Python 3.14.2是Python编程语言在2025年12月5日发布的稳定版本,属于3.14系列的第二个维护更新。该版本包含了18项修复,重点解决了多进程、数据类及正则表达式等模块的回归问题,并修复了CVE-2025-12084等安全漏洞。此版本标志着自由线程模式(移除GIL)正式获得官方支持,是Python发展的重要里程碑。
- 旧写法
checkpoint(func, x, use_reentrant=True)在 PyTorch ≥2.5 会报FutureWarning,2.6+ 可能直接抛异常 - 新写法必须写全:
checkpoint(func, x, use_reentrant=False),漏掉参数或写成True都可能引发RuntimeError: Trying to backward through the graph a second time - 如果你用的是
checkpoint_sequential,它默认就是 non-reentrant,无需额外指定
哪些层适合加 checkpoint,哪些绝对不能碰
不是所有模块都适配检查点。关键看是否“无副作用”和“确定性”:
- ✅ 推荐包裹:Transformer 的
nn.TransformerEncoderLayer、MLP 块、卷积堆叠(如 ResNet 的 bottleneck)、自注意力子模块 - ❌ 必须避开:含
nn.BatchNorm2d(统计量更新不可逆)、nn.RNN类(隐藏状态跨步依赖)、任何修改全局变量或缓存的自定义层 - ⚠️ 谨慎处理:带
torch.no_grad()或torch.set_grad_enabled(False)的局部块——检查点内禁止梯度开关,否则重算时图断裂
常见报错和绕过方法
最典型的错误是 RuntimeError: Trying to backward through the graph a second time,根本原因不是代码写错,而是检查点外又对同一张量做了 backward(),或用了共享参数(如共享 embedding 表):
- 检查是否在
checkpoint()外还调用了loss.backward()——必须只调一次,且在检查点包裹的输出上进行 - 如果模型有共享权重(如 Transformer 的 token embedding 和 lm head),确保它们不在同一个检查点函数内被多次访问;可拆成两个独立
checkpoint调用 - 调试时加
torch.autograd.set_detect_anomaly(True),能准确定位哪一行触发了图复用
真正容易被忽略的是:检查点只省激活值,不省参数、梯度或优化器状态。如果你的模型本身参数就占满显存(比如 40GB 卡跑 30GB 参数模型),加 checkpoint 也没用——得先结合 torch.compile 或 FSDP 做参数分片。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










