ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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

PyTorch 梯度检查点:用计算换显存,破解大模型训练 OOM 显存不够这件事几乎所有从单卡 demo 走向真实规模训练的人都撞过。我第一次遇到是在一台单卡机器上跑二十来层的 Transformer参数量明明不到 1Gnvidia-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、每层中间扩到 4096fp16 存储那么一层 MLP 里那个 4096 维的中间张量大概是 8 × 1024 × 4096 × 2 字节 64 MB。一个 block 里这样的中间量通常有两三个算 150 MB 一个 block24 层就是 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·an 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, expansion4, dropout0.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, vocab32000, d_model1024, n_layers24, ckpt_every6): 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, biasFalse) 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_reentrantFalse) return self.head(x)这里有几个细节值得说一下。第一_run_chunk接收了一个整数start作为参数。非张量参数传给 checkpoint 是允许的它只对张量做保存和重建其他参数原样传下去。这样写比把模块重新包成nn.Sequential更省事也不会在模块树里产生重复注册。第二use_reentrantFalse强烈建议显式写上。默认值在旧版本里是 True走的是可重入实现它对输入的要求更苛刻至少一个输入需要requires_grad也不支持 kwargs 和更复杂的嵌套结构。非重入实现是后来主推的方案torch.autograd.grad、torch.autograd.backward、以及嵌套检查点都能正常工作。第三最后一段不包检查点是个小优化。因为末尾之后没有需要重算的下游了给它加检查点等于白白多存一次边界又不省任何东西——不过这个要看你具体的分段方式有些实现里最后一段包不包差别很小实测决定。3.2 实测对比显存与耗时测量脚本本身不复杂关键是先 warmup 再 reset否则第一次运行时的 cuDNN 算法选择、内存池分配会让数据严重失真import time, torch def measure(ckpt_every, steps6, warmup2, vocab32000, d_model1024, n_layers24, bs8, seq1024, devicecuda): model ToyModel(vocab, d_model, n_layers, ckpt_every).to(device).half() opt torch.optim.AdamW(model.parameters(), lr1e-4) data torch.randint(0, vocab, (bs, seq), devicedevice) 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_noneTrue) 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 04. 我第一次上检查点踩的三个坑4.1 dropout 让重算出来的激活对不上第一次跑的时候loss 曲线比不加检查点抖得厉害步数一多还出现缓慢发散。排查了很久才想明白前向用的是 dropout mask A反向重算的时候如果又随机采样了一个 mask B那么重算出来的激活和当初前向的不是同一批数梯度自然是错的。好在torch.utils.checkpoint默认会处理这件事它会在前向时把 RNG随机数生成器状态保存下来重算前恢复。但如果你自己手写重算逻辑或者在某处手动调用torch.rand、F.dropout而没有走过去检查点的路径这个保护就失效了。我当时的代码里有一段在forward外面单独做的数据增强也有随机性恰好落在这个空洞里。经验只要涉及随机性的操作要么放在被 checkpoint 包裹的闭包内部RNG 状态会被自动保存恢复要么自己用torch.get_rng_state()/set_rng_state()手动管理。别赌。4.2 BatchNorm 的 running stats 被更新了两遍这是最阴的一个坑因为它不报错只是让准确率莫名其妙地掉一两个点。原因很简单BatchNorm 在训练模式下会更新全局的running_mean和running_var而检查点在反向时会重新跑一遍前向于是这些统计量在一次迭代里被累加了两次动量等效翻倍收敛轨迹就偏了。解决方案有三种按推荐程度排序换成 LayerNorm 或 RMSNorm。现代 Transformer 基本都是这个路线从根上绕开问题。把 BN 的momentum减半虽然有点脏但有效。冻结 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_reentrantFalse这个问题就再没出现过。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_reentrantFalse) 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_CONFexpandable_segments:True。6.2 时间开销要算有效吞吐单 step 耗时只是半个故事。真正的指标是每单位显存换来的样本吞吐也就是 samples per second per GB。检查点让显存降了 60%单步慢了 25%但 batch 从 4 提到 12有效吞吐反而涨了 1.5 倍这才是它真正的价值。测量的时候一定要做三件事warmup 至少两步。第一次迭代包含 cuDNN 算法搜索、JIT 编译、内存池首次分配时间会明显偏长。torch.cuda.synchronize()包住计时区间。CUDA 是异步的不加同步测出来的时间会是 CPU 派发核函数的时间完全是假的。固定随机种子。否则不同配置之间没法比较。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 之后才意识到的当时排查了很久最后发现是数据里混进了几条超长样本。加了个长度过滤就稳了——检查点能救显存但救不了数据里的异常值。
返回列表