商城首页欢迎来到中国正版软件门户

您的位置: 首页 > 文章列表 > 编程开发 > 如何在Python中利用PyTorch的Checkpoints技术训练超大规模网络?

如何在Python中利用PyTorch的Checkpoints技术训练超大规模网络?

  发布于2026-07-11 阅读(0)

扫一扫,手机访问

显存不够用的时候,torch.utils.checkpoint 可以说是最直接、侵入性最小的解决方案——它不需要改模型结构,不依赖多卡,也不强制分布式。核心思路很简单:用重算激活值来换显存,通常能省下 40% 到 60% 的显存,代价是训练速度会慢 20% 到 30%。

如何在Python中利用PyTorch的Checkpoints技术训练超大规模网络?

什么时候该用 checkpoint() 而不是 checkpoint_sequential()

选哪个取决于你的模型结构。如果你很清楚哪几层是“吃内存大户”,而且它们能封装成独立的函数,那就用 checkpoint(),它更灵活——可以嵌套,可以带条件逻辑,甚至支持自定义前向逻辑(比如跳过 dropout)。反过来,如果你的模型是纯 nn.Sequential 或者模块列表,想按层数粗粒度切分(比如每 4 层一组),那 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 状态没对齐。举个例子,下面这种写法就是典型的坑:

  • 错误写法:def block(x): return self.norm(x) + self.ffn(x) —— 这里的 self 是闭包变量,反向重算时拿不到当前 self.norm 的状态。正确的做法是显式传参:def block(x, norm, ffn): return norm(x) + ffn(x)
  • 另一个常见遗漏:函数内如果有 nn.Dropout,必须确保 preserve_rng_state=True(默认就是 True),否则重算时 dropout 行为不一致。千万不要手动关掉它。
  • 还有输入张量的问题: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。反过来,如果先对原始模型加 checkpoint 再喂给 FSDP,FSDP 初始化阶段会因为计算图未就绪而失败。

和 DDP 混用没有硬性冲突,但要注意 DDP 的 find_unused_parameters=True 可能会和 checkpoint 内部的梯度路径检测冲突。建议关掉这个参数,改用 Hugging Face Transformers 封装的 gradient_checkpointing_enable()。另外,一个性能提醒:checkpoint + FSDP 会放大通信等待时间,因为重算激活期间 GPU 在空转。可以配合 backward_prefetch=BackwardPrefetch.BACKWARD_PRE 来缓解。

真正容易被忽略的点

最后说一个很多人容易误解的地方:checkpoint 只节省中间激活的显存,不节省权重本身的显存。如果你的模型参数本身就超了显存(比如 7B 模型放不进 24G 卡),那它救不了你——这时候必须上 FSDP、TP 或者量化。另外,torch.compilecheckpoint 默认不兼容,除非显式启用 use_reentrant=False 并禁用部分优化(比如 dynamic=True 会报错)。这些细节如果不注意,调试起来会很头疼。

本文转载于:https://www.php.cn/faq/2799649.html 如有侵犯,请联系zhengruancom@outlook.com删除。
免责声明:正软商城发布此文仅为传递信息,不代表正软商城认同其观点或证实其描述。

热门关注