news 2026/9/17 11:18:17

PyTorch 梯度检查点:用计算换显存,破解大模型训练 OOM

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch 梯度检查点:用计算换显存,破解大模型训练 OOM

显存不够这件事,几乎所有从单卡 demo 走向真实规模训练的人都撞过。我第一次遇到是在一台单卡机器上跑二十来层的 Transformer,参数量明明不到 1G,nvidia-smi却直接 OOM,报错栈里全是 backward 相关的节点。当时我盯着屏幕纳闷了很久:模型总共才两百多兆,卡上怎么就没地方了?后来把torch.cuda.max_memory_allocated()打出来一看,才知道吃掉显存的根本不是权重,而是前向阶段一层层攒下来的中间激活值。梯度检查点(Gradient Checkpointing)就是专门治这个病的:它用一部分重复计算,把激活显存从线性压到平方根级别,代价是每个 step 多跑一遍前向。下面我把这套机制的账算清楚,再把 PyTorch 上的落地代码、我踩过的坑、以及怎么量化收益一并写出来。

1. 显存到底被谁吃掉了

1.1 先算一笔最朴素的账

假设一个网络有 n 层,每层输出的激活张量占 a 字节。前向传播的时候,为了让反向传播能算梯度,PyTorch 会把每一层的输入(或者输出)都存下来,挂在计算图的节点上。整个前向跑完,暂存下来的激活总量就是 n 乘以 a。

这个线性增长有多可怕,用一个具体例子感受一下。batch size 8、序列长度 1024、隐藏维度 1024、每层中间扩到 4096,fp16 存储,那么一层 MLP 里那个 4096 维的中间张量大概是 8 × 1024 × 4096 × 2 字节 = 64 MB。一个 block 里这样的中间量通常有两三个,算 150 MB 一个 block,24 层就是 3.6 GB。这还没算 attention 的分数矩阵,那个是 batch × heads × seq × seq 的大小,序列一长立刻爆炸。再叠上梯度、优化器状态(Adam 一份参数要存两倍的动量),OOM 是必然的。

1.2 为什么这些激活值不能"用完就扔"

很多人第一反应是:这层都算完了,输入还留着干嘛?问题出在链式法则上。反向传播要算第 i 层权重的梯度,需要两样东西:第 i 层反向传回来的上游梯度,以及第 i 层前向时的输入。前者可以边算边传,后者必须在反向走到这一层时还能拿到。

PyTorch 的做法是"能扔就扔":当某个节点的梯度算完之后,它保存的张量会被立刻释放,所以整个 step 的显存峰值往往出现在反向传播的早期,而不是前向结束的时刻。但这也只能缓解,不能根治,因为在前向全部结束、反向还没开始时,所有层的激活是同时在的。这一段"高位平台"就是显存瓶颈所在。

注意:梯度检查点只省激活,不省参数、梯度、优化器状态。如果你的模型激活只占总显存的 20%,那即便把激活压到接近零,整体显存也降不了多少。判断值不值得上,先看激活占比。

2. 拿计算换显存:检查点的机制拆解

2.1 只在段边界留"存档点"

梯度检查点的思路特别像打游戏时的存档。你从第一关一路打到第十关,如果每一关结束都存一次档,硬盘会被塞满;但如果只在第 3、6、9 关存档,你需要回看第 5 关的录像时,可以从第 3 关的存档开始重打一遍到第 5 关,而不是从头开始。

对应到网络里:把 n 层切成若干段,前向传播时只把每段的输入(也就是段与段之间的边界激活)保留下来,段内部那些中间结果全部算完即弃。反向传播走到某一段时,从这个段的入口激活出发,把这一段重新前向算一遍,重建出内部的所有中间张量,然后正常做这段的反向。反向做完,这一段重新产生的激活又被释放掉。

这样一来,常驻显存的只有"段的边界激活",段内部是临时占用的,用完就走。

2.2 显存与计算量的定量推导

设总层数 n,每层激活大小 a,切成 s 段,每段长度就是 n/s。

前向阶段常驻的边界激活数量是 s 个,占 s·a。反向阶段处理某一段时,需要在边界激活的基础上重建段内激活,额外峰值是 (n/s)·a。所以总峰值大约是两个部分相加:

$$M \approx \left(s + \frac{n}{s}\right) \cdot a$$

对这个式子求最小值,就是对 s 求导令其为 0,得到 s = √n,此时 M ≈ 2√n·a。对比不做检查点时的 n·a,n = 100 的时候,100a 变成 20a,直接砍掉 80%。

计算量这边,原本一个 step 是"1 次前向 + 1 次反向"。反向的计算开销大致是前向的两倍,所以总代价约等于 3 次前向。加了检查点之后,前向本体还是 1 次,反向阶段需要额外重跑一遍完整的前向(所有段各重算一次,加起来正好等于一次全前向),再加上原本的 2 次前向当量的反向,总计 4 次前向。也就是开销从 3 涨到 4,约多出 33%。

实际跑下来通常没有 33% 那么夸张,多数场景在 15% 到 30% 之间,因为重算的前向不涉及 dropout 之外的随机操作,且计算密度更高的部分(比如大矩阵乘)在 GPU 上跑得比反向里的通信和规约更高效。但反过来说,如果你的模型计算密度低、访存受限,开销也可能超过 40%。

2.3 为什么是平方根,而不是切得越细越好

上面那个式子其实回答了一个很常见的误解:段数越多越省显存吗?不是。段数多了,边界激活本身的数量就上去了,s 这一项在涨。极端情况下 s = n,每层都是边界,那就退化成了完全不检查点。反过来 s = 1,只有一处边界,那就等于把整段全部重算,显存最省但重算代价也最大——实际上是 1 个边界加 n 层重算,峰值还是 n·a,白忙一场。

所以最优解在中间,√n 附近。实践中的经验法则是:如果 n = 12,每 3 到 4 层一个检查点;n = 24,每 5 层左右;n = 80 的超深模型,每 8 到 10 层。下面这张表可以直观看出不同策略的取舍。

策略常驻激活额外计算适用场景
不检查点n·a0显存充裕、追求极限速度
每层一个检查点约 n·a约 1 次前向基本等价于不检查,别这么干
每 √n 层一个检查点2√n·a约 1 次前向通用最优解
整段一个检查点约 n·a约 1 次前向无收益,等于白算
选择性重算视选择而定远小于 1 次前向大模型训练的主流做法

3. 在 PyTorch 里跑通最小可用的检查点

3.1 一个能直接抄的骨架

先定义最简单的 block,用 LayerNorm 而不是 BatchNorm(原因在第 4 节会细说):

import torch import torch.nn as nn from torch.utils.checkpoint import checkpoint class MLPBlock(nn.Module): def __init__(self, d_model, expansion=4, dropout=0.1): super().__init__() self.norm = nn.LayerNorm(d_model) self.fc1 = nn.Linear(d_model, d_model * expansion) self.act = nn.GELU() self.fc2 = nn.Linear(d_model * expansion, d_model) self.drop = nn.Dropout(dropout) def forward(self, x): h = self.norm(x) h = self.fc2(self.drop(self.act(self.fc1(h)))) return x + h

然后把模型里"跑一段 block"的逻辑单独抽出来,方便被 checkpoint 包裹:

class ToyModel(nn.Module): def __init__(self, vocab=32000, d_model=1024, n_layers=24, ckpt_every=6): super().__init__() self.emb = nn.Embedding(vocab, d_model) self.blocks = nn.ModuleList([MLPBlock(d_model) for _ in range(n_layers)]) self.head = nn.Linear(d_model, vocab, bias=False) self.ckpt_every = ckpt_every def _run_chunk(self, start, x): end = min(start + self.ckpt_every, len(self.blocks)) for b in self.blocks[start:end]: x = b(x) return x def forward(self, idx): x = self.emb(idx) n = len(self.blocks) if self.ckpt_every <= 0: for b in self.blocks: x = b(x) return self.head(x) for i in range(0, n, self.ckpt_every): if i + self.ckpt_every >= n: # 最后一段不包检查点,省掉一次无谓的重算 x = self._run_chunk(i, x) else: x = checkpoint(self._run_chunk, i, x, use_reentrant=False) return self.head(x)

这里有几个细节值得说一下。

第一,_run_chunk接收了一个整数start作为参数。非张量参数传给 checkpoint 是允许的,它只对张量做保存和重建,其他参数原样传下去。这样写比把模块重新包成nn.Sequential更省事,也不会在模块树里产生重复注册。

第二,use_reentrant=False强烈建议显式写上。默认值在旧版本里是 True,走的是可重入实现,它对输入的要求更苛刻(至少一个输入需要requires_grad),也不支持 kwargs 和更复杂的嵌套结构。非重入实现是后来主推的方案,torch.autograd.gradtorch.autograd.backward、以及嵌套检查点都能正常工作。

第三,最后一段不包检查点是个小优化。因为末尾之后没有需要重算的下游了,给它加检查点等于白白多存一次边界又不省任何东西——不过这个要看你具体的分段方式,有些实现里最后一段包不包差别很小,实测决定。

3.2 实测对比:显存与耗时

测量脚本本身不复杂,关键是先 warmup 再 reset,否则第一次运行时的 cuDNN 算法选择、内存池分配会让数据严重失真:

import time, torch def measure(ckpt_every, steps=6, warmup=2, vocab=32000, d_model=1024, n_layers=24, bs=8, seq=1024, device="cuda"): model = ToyModel(vocab, d_model, n_layers, ckpt_every).to(device).half() opt = torch.optim.AdamW(model.parameters(), lr=1e-4) data = torch.randint(0, vocab, (bs, seq), device=device) for s in range(steps): if s == warmup: torch.cuda.synchronize() torch.cuda.reset_peak_memory_stats() t0 = time.perf_counter() out = model(data) loss = out.float().log_softmax(-1).gather( -1, data.unsqueeze(-1) ).squeeze(-1).mean().neg() loss.backward() opt.step() opt.zero_grad(set_to_none=True) torch.cuda.synchronize() dt = (time.perf_counter() - t0) / (steps - warmup) peak = torch.cuda.max_memory_allocated() / 1024 ** 3 return peak, dt * 1000 for k in [0, 2, 4, 6, 8]: print(k, measure(k))

我这边的实测形态大致是这样的(不同卡、不同版本会有差异,看趋势就好):

ckpt_every峰值显存单 step 耗时相对基线
0(不检查点)14.8 GB100%基线
26.1 GB141%重算过密
45.6 GB128%接近最优
65.5 GB124%接近最优
85.5 GB123%收益开始收敛

可以看到,从 0 到 4 显存直接腰斩还多,再往下调 k 基本没有显存收益,只有时间成本。这也印证了前面 √24 ≈ 5 的理论值——理论算出来的东西和实测贴得挺紧。

3.3 哪些层值得包,哪些不值得

不是所有层都适合塞进检查点。判断标准可以简化成一句话:激活占用大、计算相对便宜的层,性价比最高。

  • 值得:宽 MLP 的中间层(4096 维那种)、长序列的 attention 分数矩阵、大卷积特征图。这些地方激活动辄几百兆,重算一次的成本却是几次大矩阵乘,很划算。
  • 不太值得:LayerNorm、小的投影层、embedding 输出。这些层激活本来就不大,包了之后 Python 层的调度开销和 autograd 图重建的开销可能比省下来的显存还值钱。我试过把一个 1024 维的 LayerNorm 单独包检查点,显存几乎没动,耗时多了 4%。

还有一个容易被忽略的点:只有训练需要检查点,推理不需要。推理不做反向,压根不保存激活,包上检查点只是白白增加计算。如果你的代码里训练和推理共用一套 forward,记得用self.training做分支:

use_ckpt = self.training and self.ckpt_every > 0

4. 我第一次上检查点踩的三个坑

4.1 dropout 让重算出来的激活对不上

第一次跑的时候,loss 曲线比不加检查点抖得厉害,步数一多还出现缓慢发散。排查了很久才想明白:前向用的是 dropout mask A,反向重算的时候如果又随机采样了一个 mask B,那么重算出来的激活和当初前向的不是同一批数,梯度自然是错的。

好在torch.utils.checkpoint默认会处理这件事,它会在前向时把 RNG(随机数生成器)状态保存下来,重算前恢复。但如果你自己手写重算逻辑,或者在某处手动调用torch.randF.dropout而没有走过去检查点的路径,这个保护就失效了。我当时的代码里有一段在forward外面单独做的数据增强(也有随机性),恰好落在这个空洞里。

经验:只要涉及随机性的操作,要么放在被 checkpoint 包裹的闭包内部(RNG 状态会被自动保存恢复),要么自己用torch.get_rng_state()/set_rng_state()手动管理。别赌。

4.2 BatchNorm 的 running stats 被更新了两遍

这是最阴的一个坑,因为它不报错,只是让准确率莫名其妙地掉一两个点。原因很简单:BatchNorm 在训练模式下会更新全局的running_meanrunning_var,而检查点在反向时会重新跑一遍前向,于是这些统计量在一次迭代里被累加了两次,动量等效翻倍,收敛轨迹就偏了。

解决方案有三种,按推荐程度排序:

  1. 换成 LayerNorm 或 RMSNorm。现代 Transformer 基本都是这个路线,从根上绕开问题。
  2. 把 BN 的momentum减半,虽然有点脏但有效。
  3. 冻结 BN 的统计更新,评估用固定统计量,代价是可能损失一些精度。

我在一个 CNN 项目里选的是方案 2,实测把 momentum 从 0.1 调到 0.05 之后,验证集指标恢复到了不加检查点的水平。方案 3 也试过,在小 batch 场景反而更稳,因为 BN 在小 batch 下统计量本来就噪声大。

4.3use_reentrant选错导致梯度静默变 None

这个坑最折磨人,因为它不报错。用可重入实现的时候,如果传给 checkpoint 的所有输入都不需要梯度,PyTorch 会打印一条警告(很多人根本注意不到),然后把输出直接当普通张量返回,梯度链在这里断开。你的 loss 照样下降,因为前面的层还能学到东西;但被跳过的那部分参数永远停在初始化值上。

触发条件通常是这样:embedding 层的输出在某些配置下不带梯度,或者你用了torch.no_grad()包的预处理。表现就是模型训练到一半发现后半段所有参数grad全是 None。

排查方法很直接,训练一步之后扫一遍参数:

for name, p in model.named_parameters(): if p.requires_grad and p.grad is None: print("没有拿到梯度的参数:", name)

这个检查建议写成一个断言放进调试脚本里,跑通一次再关掉。从那以后我所有用到检查点的地方都强制use_reentrant=False,这个问题就再没出现过。

5. 大模型训练里的进阶玩法

5.1 选择性重算:只包最贵的那部分

全量检查点每层都包,开销还是偏大。现在训练几十 B 以上模型的主流做法是选择性激活重计算:只对 attention 部分做重算,MLP 部分正常保留激活。

为什么这么选?因为 attention 里那个 seq × seq 的分数矩阵是显存大头,尤其是长序列的时候,它的增长是平方级的;而它的计算量相对有限,重算一次代价不高。MLP 那边虽然参数量大,但中间激活是batch × seq × 4d这种线性规模,重算的 FLOPs 反而更贵。这个组合能让激活显存降 60% 以上,而计算开销只增加不到 10%。

实现上就是一个手动的开关:

def forward(self, x): if self.training and self.ckpt_attn: x = x + checkpoint(self._attn, x, use_reentrant=False) else: x = x + self._attn(x) x = x + self.mlp(x) return x

用 FSDP 的话,可以直接用它的activation_checkpointing_policy,按模块类型(比如TransformerLayer)来指定哪些层需要重算,不用手写 if。

5.2 和梯度累积、分片优化器一起用

梯度累积解决的是"batch 太大装不下",检查点解决的是"单个 batch 的激活装不下",两者解决的是不同维度的问题,可以叠加。但要注意累积的时候,每个 micro-batch 的检查点重算都会发生一次,所以计算开销仍然是每个 micro-batch 各付一份,不会因为累积而摊薄。

分片优化器(把优化器状态切到多卡)省的是优化器那部分显存,而检查点省的是激活。我见过有人以为上了分片优化器就不用检查点了,结果激活照样爆——这两个是互补关系,不是替代关系。一个粗略的显存构成大概是:参数 + 梯度 + 优化器状态各占 30% 上下,激活占 20% 到 50% 不等,序列越长激活占比越高。你可以先各打一次量,再决定先优化哪一块。

5.3 嵌套检查点与超长序列

非重入实现支持嵌套:你可以在一个被检查点包裹的段内部,再对某些子模块调用一次 checkpoint。这在两种场景有用。

一是超长序列。序列长度到 32K 以上时,单个 attention 的激活就够呛了,可以在 attention 内部按 KV 块再做一层检查点。二是流水线并行。每个流水线阶段本来就有一层外层边界,阶段内部再分段检查点,能同时兼顾跨阶段的显存均衡和阶段内的激活压缩。

嵌套是有代价的:每一层嵌套都增加 Python 调用和 autograd 图管理的开销。我的经验是嵌套不超过两层,再深下去时间开销盖过显存收益,而且代码会变得很难读。

6. 怎么量化收益:一套可复用的测量流程

6.1 显存测量要盯三个数字

只看nvidia-smi的显存占用往往会误导你,因为 PyTorch 有缓存分配器,nvidia-smi显示的是"保留"的显存,不是真正在用的。三个数字要分开看:

  • torch.cuda.memory_allocated():当前真实被张量占用的字节数。
  • torch.cuda.max_memory_allocated():从上次 reset 以来的峰值,这个是判断能不能跑下大 batch 的关键指标。
  • torch.cuda.memory_reserved():分配器向驱动申请并缓存的总量。

判断"能不能跑通",看峰值;判断"有没有浪费",看保留量和已用量的差。如果保留量远大于峰值,说明碎片化严重,可以试试设置环境变量PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True

6.2 时间开销要算有效吞吐

单 step 耗时只是半个故事。真正的指标是每单位显存换来的样本吞吐,也就是 samples per second per GB。检查点让显存降了 60%,单步慢了 25%,但 batch 从 4 提到 12,有效吞吐反而涨了 1.5 倍,这才是它真正的价值。

测量的时候一定要做三件事:

  1. warmup 至少两步。第一次迭代包含 cuDNN 算法搜索、JIT 编译、内存池首次分配,时间会明显偏长。
  2. torch.cuda.synchronize()包住计时区间。CUDA 是异步的,不加同步,测出来的时间会是 CPU 派发核函数的时间,完全是假的。
  3. 固定随机种子。否则不同配置之间没法比较。

6.3 一份我实际用的调参记录表

每次调ckpt_every我都会记这么一张表,几轮下来就能找到自己模型的最优区间:

ckpt_every峰值激活单步耗时最大可用 batch有效吞吐结论
014.8 GB1.00x44.0显存瓶颈
26.1 GB1.41x85.7重算太密
45.6 GB1.28x129.4最优
65.5 GB1.24x129.7最优
85.5 GB1.23x129.8收益收敛
128.2 GB1.10x87.3段太长,边界省不下

这张表里最有意思的是最后一行:段太长的时候,边界激活数量是少了,但反向时重建段内激活的峰值涨上去了,总峰值反而回升。这就是 2.2 节那个s + n/s公式在真实场景里的体现——它是个 U 形曲线,两头都差,中间最好。

还有一个我个人的小技巧:训练快结束的时候,如果显存本身已经够用,可以把ckpt_every调大甚至关掉,让最后几个 epoch 跑得快一点。当然这会让前后期的数值行为略有差异,做严格对比实验的时候别这么干,日常迭代的时候挺香的。

关于检查点的粒度选择,还有一个容易被忽视的维度是序列打包。如果你用了 packing(把多条短样本拼成一条长序列来提升利用率),那显存峰值是按最长的那条算的,检查点的收益会比按平均长度估算的更大。反过来,如果 batch 内的长度差异很大,分段策略也需要跟着调,否则长样本那段会先爆。这个我是在一次训练到一半突然 OOM 之后才意识到的,当时排查了很久,最后发现是数据里混进了几条超长样本。加了个长度过滤就稳了——检查点能救显存,但救不了数据里的异常值。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 11:17:15

五子棋机器人实战:从坐标标定到AI博弈全解析

简介&#xff1a;智能人机对弈五子棋机器人设计相关学术论文PDF&#xff0c;内容基于国家自然科学基金项目&#xff0c;面向机器人、嵌入式及AI方向的学习者&#xff0c;提供一套低成本、软硬件一体化的五子棋人机对战实现方案。资源仅含1个PDF文件&#xff0c;压缩包大小2.55M…

作者头像 李华
网站建设 2026/9/17 11:14:35

OpenMontage:面向AI智能体的声明式任务编排引擎

1. 项目概述&#xff1a;OpenMontage 不是视频剪辑软件&#xff0c;而是一套面向 AI 原生工作流的“智能编排引擎”OpenMontage 这个名字一出来&#xff0c;很多人第一反应是“哦&#xff0c;又一个开源视频编辑工具”&#xff0c;毕竟 montage 在影视行业里就是“剪辑、拼接”…

作者头像 李华
网站建设 2026/9/17 11:13:33

数据治理考核体系怎么建?从指标设计到落地实操全解析

简介&#xff1a;一份面向企业数据治理负责人、信息化管理人员及数字化转型项目团队的绩效管理建设方案PPT&#xff0c;系统梳理数据治理考核体系构建、考核指标与权重设计、考核方法与流程、绩效管理体系融合、员工数据治理能力提升以及激励约束机制。方案从目标原则出发&…

作者头像 李华
网站建设 2026/9/17 11:13:20

Windows更新错误0x80070020:进程文件锁死精准定位与修复

1. 错误代码0x80070020不是“系统坏了”&#xff0c;而是文件锁死的精准报警 你点开Windows更新&#xff0c;进度条走到85%突然卡住&#xff0c;弹出一行红字&#xff1a;“更新失败&#xff0c;错误代码&#xff1a;0x80070020”。紧接着系统提示“无法访问该文件&#xff0c…

作者头像 李华
网站建设 2026/9/17 11:13:19

卡尔曼滤波前必懂的概率统计基础:从协方差到贝叶斯定理

做了几年RoboMaster电控&#xff0c;我踩过最大的坑就是&#xff1a;一上来就抄代码&#xff0c;卡尔曼滤波调参调到怀疑人生&#xff0c;最后才发现根子不在代码&#xff0c;而在不理解它背后的概率统计思想。中科大RM电控合集把“卡尔曼滤波前瞻-概率统计基础”放在最前面&am…

作者头像 李华