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

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() 上下文——如果是,那就会出问题。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.compile 和 checkpoint 默认不兼容,除非显式启用 use_reentrant=False 并禁用部分优化(比如 dynamic=True 会报错)。这些细节如果不注意,调试起来会很头疼。
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
正版软件
正版软件
正版软件
正版软件
正版软件
1
2
3
7
8