ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch显存管理揭秘:从Autograd到梯度检查点的优化实战

PyTorch显存管理揭秘:从Autograd到梯度检查点的优化实战 1. 先解决那个所有炼丹师都骂过的OOM为什么显存知识值得单独开一讲干深度学习这一行最扫兴的事情莫过于训练跑到一半CUDA out of memory直接砸脸上。模型前向算得好好的数据也在 GPU 上躺着可loss.backward()一调用显存瞬间爆掉。这种问题新手第一反应是调小 batch size老手会先看一眼是不是哪个中间变量没释放但真正能一句话说清PyTorch 的显存到底是被谁吃掉的的人其实不多。这一讲核心就三件事PyTorch 的动态 DAG有向无环图是怎么搭起来的、Autograd 反向求导在图上到底怎么跑、以及被称为 Activation 的中间激活值为什么是显存管理的大头。这三件事串起来你才能准确回答一个看似简单的问题——训练一个模型时显存从什么时候开始涨又是在什么时候被释放的理解这套机制最直接的红利就是你以后处理 OOM、跑大模型、做梯度检查点、调 batch size 的时候不再靠试错而是能算着显存写代码。这一讲的内容不是那种会用 fit() 就行的层面而是深入到框架内部的运作逻辑适合已经开始写自定义训练循环、打算在显存上做文章的人。如果你刚入门也建议硬着头皮读完因为 PyTorch 里坑最多的几个操作——detach()、retain_graph、inplace修改——全都在这一讲的范围里。我先把结论摆在前面前向传播的过程就是动态建图的过程反向传播的过程就是按图求导并把用完的缓存丢掉的过_程而 Activation 之所以占显存恰恰是因为它需要留着给反向用的“中间变量”。后面所有内容都是围绕这句大白话展开的。2. 前向传播悄悄做的三件事Tensor、DAG 与 grad_fn 的诞生2.1 一个 Tensor 的三要素90% 的人只用了两个每个 PyTorch 的Tensor对象表面上看就是一个带shape、dtype、device的数组但当你打开requires_gradTrue这个开关之后它的内部结构瞬间复杂了一个维度。从 Autograd 的角度看一个 Tensor 至少包含三样东西data真实的数值数据也就是存储张量内容的内存空间requires_grad是否需要在该张量参与运算时追踪梯度grad_fn记录这个张量是怎么算出来的的函数对象。新手最常见的误区是把注意力全放在data上觉得深度学习就是张量运算而忽略了grad_fn才是 Autograd 的灵魂。grad_fn不是一个简单的属性标签它是一个持有反向传播方法的函数节点。你执行y x * w b的时候得到的y的grad_fn会指向一个乘法或加法的反向节点这个节点记录了参与运算的输入 tensor 的引用。正因如此前向过程结束后你手上拿到的每一个中间张量都自带一条我从哪来的完整记录。可以做个简单实验来验证。定义两个叶子张量做一次组合运算然后打印中间结果的grad_fnimport torch x torch.randn(3, 3, requires_gradTrue) w torch.randn(3, 3, requires_gradTrue) y x * w z y.sum() print(x.grad_fn) # None叶子节点没有 grad_fn print(y.grad_fn) # MulBackward0 object at ... print(z.grad_fn) # SumBackward0 object at ...看到没x是用户手动创建的叶子节点它的grad_fn是None但y和z的grad_fn分别是MulBackward0和SumBackward0。这条链子一旦形成反向传播就不用你再手动写链式法则了框架自己就能顺藤摸瓜。2.2 动态 DAG 到底动态在哪每次前向都重新搭图PyTorch 使用的计算图模型是动态 DAG这跟 TensorFlow 1.x 时代的静态图模型有本质区别。静态图是先画图再喂数据你定义网络结构时框架就已经把整张计算图确定下来了后续所有 batch 都在这张固定的图上执行。PyTorch 的做法完全不同每次执行前向传播都会重新构建一张全新的计算图。这意味着什么首先你的网络结构可以在运行时动态变化。if条件、for循环、递归调用这些 Python 原生控制流可以直接写进模型里框架不需要预先编译。这正是 PyTorch 在研究和调试阶段碾压静态图框架的核心原因——你可以像一个普通 Python 程序一样去调试你的神经网络。但动态图也有代价。因为每次前向都要现场搭图边搭边保存中间结果所以它相比静态图会有一些额外的内存开销和调度开销。这也是为什么 PyTorch 后来推出了torch.compile和torch.jit.script来尝试把动态图静态化一部分从而获得性能提升。动态 DAG 还有一个容易忽略的特性图的方向是从数据到结果的而反向求导的方向是从结果到数据的。前向搭图时每执行一行代码就往图上追加一个新节点。这个过程不可回退除非被detach切断了连接。因此前向过程中显存会一直累积增长这个累积的正是我们后面要重点讲的 Activation 以及计算图节点本身。2.3 叶子节点为什么特殊一个关乎梯度生死的重要概念在 Autograd 体系里叶子节点leaf tensor指的是由用户直接创建、不依赖任何其他张量运算得到的 Tensor。它有几个非常关键的特质只有叶子节点的grad属性会在反向传播时被自动填充。非叶子节点的grad默认在反向计算完之后就被清空了除非你调用retain_grad()强制保留。优化器的step()方法只遍历model.parameters()而模型参数的requires_gradTrue且是叶子节点所以它们的.grad会被填充并用于参数更新。叶子节点在创建时如果requires_gradTrue它的grad_fn永远是None因为它不是通过任何运算生成的。第二条非常重要它直接解释了为什么你在训练循环里能访问param.grad来手动更新参数而中间变量的梯度却拿不到。框架这么做不是小气而是显存策略的一部分非叶子节点的梯度只在反向传播过程中临时存在算完就丢绝不长期占用显存。后面我会专门展开讲这个设计背后的显存账本。3. 一块 GPU 显存从分配到释放的完整生命周期谁在涨、谁在跌、峰值在哪3.1 前向阶段显存像滚雪球一样涨上去现在我们把目光聚焦到显存。很多人以为显存主要是被模型参数占掉的实际上在一个典型的大模型训练任务里参数本身只是很小的一部分。我们来掰着手指头算一笔账。假设你有一个参数量为 N 的模型以 FP16 精度存储参数那参数本身只占2 × N字节。但反向传播需要计算梯度梯度的精度必须足够高你至少还得准备一份 FP32 或者 FP16 的梯度张量这又是2 × N或4 × N字节。再加上优化器状态——如果是 Adam 优化器需要额外保存一阶动量 m 和二阶动量 v这又是4 × N或8 × N字节。把这些加起来参数量为 N 的模型光参数、梯度、优化器状态这三件套通常就要吃掉12 × N到20 × N字节的显存。七B 参数的大模型光这套基础开销就在 100GB 以上所以大家才会去研究 LoRA、量化这些省显存的技术。但这还只是基础开销。前向传播开始后真正让显存失控的是中间激活值。看下面这段代码def forward(self, x): h1 torch.relu(self.fc1(x)) # 中间张量 h1 h2 torch.relu(self.fc2(h1)) # 中间张量 h2 out self.fc3(h2) # 输出 out return out在 PyTorch 的默认设置下h1、h2这些中间结果都会因为被后续操作引用而保存在显存里目的只有一个等反向传播时Autograd 需要用到它们来计算梯度。保存 Activation 的显存开销正比于 batch size × 序列长度 × 隐藏维度 × 网络层数。对一个 12 层的 Transformer 来说单个 batch 的激活值大小经常是模型参数的几倍甚至十几倍。这就导致了一个反直觉的现象你的显存主要不是被模型装掉的而是被模型算掉的。很多人费尽心思把模型参数从 FP32 换成 FP16显存却依然紧张原因就是他们没动激活值这块最肥的肉。3.2 反向阶段Autograd 引擎的随算随扔反向传播开始后显存的走势和前向刚好相反——一路向下释放。但释放的时机和粒度很讲究。loss.backward()被调用后PyTorch 的 Autograd 引擎会从loss这个标量出发沿着grad_fn链条逆着 DAG 的方向遍历。每经过一个节点它要用前向保存的激活值来计算局部梯度然后把梯度传给上游节点最后更新到叶子节点的.grad属性上。关键点来了每个节点一旦完成自己的梯度计算它保存的前向激活值缓存如果不再被其他节点需要就会被立即释放。也就是说反向传播是一个随算随扔的过程。这就是为什么你在跑反向传播的时候用nvidia-smi观察显存会看到一个从峰值逐渐下降的曲线。这个设计其实非常优雅。如果 PyTorch 在前向结束时一次性把所有激活值都留着反向结束后再统一释放那显存峰值会更高而且释放逻辑也更粗放。按需计算、按需释放才能让显存在整个训练循环中保持一个相对稳定的水位线。3.3 显存峰值为什么 OOM 总在前向结束、反向还没开始的瞬间理解了前向分配和反向释放的机制你就能回答一个经典问题显存峰值出现在什么时候答案很明确出现在前向传播刚结束、反向传播还未开始的临界点。这一刻模型参数、优化器状态、梯度缓冲全部就位而且所有层的激活值都还完整保存在显存里。一旦反向开始激活值会一层层释放显存压力随之缓解。所以 OOM 往往发生在刚刚调用backward()的时候——不是backward()本身需要额外开多少显存而是backward()被调用的瞬间前向刚结束所有临时缓存还满满当当此刻恰好是整个训练周期里显存需求的最高峰。基于这个结论你可以得出两个优化思路要么减小前向过程的峰值需求比如减小 batch size、做梯度累积要么让一部分激活值不要在峰值窗口内占着显存比如用梯度检查点技术把前向过程切成若干段重算。这两个思路正好对应了后面第四、五节的内容。4. 反向求导的黑盒拆解DAG 上的链式法则怎么一步步跑4.1 backward() 到底在遍历什么一张自带路径的计算图前面说了前向构建的每个 Tensor 都携带grad_fn这个grad_fn对象内部又持有指向输入张量的引用而输入张量又有自己的grad_fn。这样一层层嵌套就织成了一张巨大的反向传播网络。当我们调用loss.backward()时PyTorch 做的工作本质上就是在这张网络上做一次反向拓扑排序遍历。从loss这个输出节点出发沿着每个节点的grad_fn去找到它的输入节点逐层往上游传播。每个节点在遍历过程中都会利用前向保存的输入激活值计算出当前节点的输出对输入的局部雅可比再乘以上游传来的梯度得到传往更上游的梯度信号。这个过程的数学基础就是链式法则。简单来说如果z f(y)且y g(x)那么dz/dx dz/dy × dy/dx。Autograd 并不需要你手动推导这个公式它只是把这个计算过程机械化地作用在 DAG 的每一条边上。这里可以用一个极简例子来演示。定义x 2然后x torch.tensor(2.0, requires_gradTrue) y x * 3 # dy/dx 3 z y ** 2 # dz/dy 2y 12链式法则: dz/dx 3 * 12 36 z.backward() print(x.grad) # tensor(36.)backward()执行时Autograd 引擎的工作方式是先算出dz/dz 1然后从z节点传播到y节点乘以dz/dy 12得到传到y处的梯度12再沿着y的grad_fn传播到x乘以dy/dx 3最终得到x.grad 36。每一步的局部导数都是前向保存好的根本不需要重新算一遍前向就能拿到。4.2 多路径汇合同一 Tensor 被多处使用梯度怎么求和DAG 之所以叫图而不是树是因为一个节点的输出可能会被多个下游节点使用。举个例子x torch.randn(4, requires_gradTrue) a x * 2 b x * 3 loss a.sum() b.sum() loss.backward()这里的x有两个分支a分支和b分支。根据链式法则dloss/dx应该是两条路径梯度之和。PyTorch 处理这种情况的方式是逐个遍历所有路径把梯度累加到同一个x.grad缓冲里。这也是为什么.grad属性是一个累加值而不是覆盖值——它天然支持多路径梯度累加。这个特性还有一个实际应用梯度累积gradient accumulation。当你的 batch size 太大放不进显存时可以把一个 batch 拆成多个 mini-batch分别前向和反向然后把梯度累加足够的步数后再调用优化器的step()。由于.grad本来就是累加语义这个操作只需要小心地控制zero_grad()的时机即可不需要任何额外的机制配合。4.3 非叶子节点梯度默认不保留一次标准的显存账本算计前面提到非叶子节点的.grad在反向传播完成后会被清空这是 PyTorch 内存池设计里非常精妙的一笔。设想一下如果你有 100 层网络每层都有若干中间激活值。如果不做任何清理反向传播结束后每个非叶子节点的.grad都会留在显存里。这些梯度数量级和激活值接近但它们在参数更新这一环节又完全用不上——优化器只关心叶子节点即模型参数的梯度。留在那里不仅浪费显存还会让下一次前向的显存分配变得更加紧张。所以 PyTorch 默认在反向传播时每算完一个非叶子节点的梯度用完后就直接丢弃只保留叶子节点的.grad。如果你确实需要某个中间变量的梯度来做梯度裁剪、特征可视化、或者调试你可以在前向过程中对该张量调用retain_grad()方法明确告诉框架这个节点的梯度请帮我留着。来看一个带retain_grad()的示例x torch.randn(3, 3, requires_gradTrue) y x ** 2 z y.mean() y.retain_grad() # 强制保留 y 的梯度 z.backward() print(x.grad) # tensor([[0.6667, ...]]) 或类似正确梯度 print(y.grad) # tensor([[0.3333, ...]])被强制保留了下来如果不加y.retain_grad()y.grad在z.backward()后会变成None。这个细节测试模型中间层梯度时非常实用但不建议在训练循环里对每一层都这么干因为这会重蹈显存爆炸的覆辙。4.4 原地操作的版本计数器为什么 inplace 修改会直接报错这是 Autograd 中最让新手头疼的问题之一为什么x 1或者y.relu_()这类原地操作在 PyTorch 的反向计算图里经常会抛出RuntimeError: a leaf Variable that requires grad is used in an in-place operation。原因很简单。Autograd 保存的是引用在前向时它会记住每个参与运算的 Tensor 对象。原地操作直接修改 Tensor 的数值内容而不是创建一个新的 Tensor这相当于在前向传播结束后把已经登记在计算图里的输入数据悄悄换掉了。等到反向传播时Autograd 拿它之前保存的数值去算局部导数会发现数值对不上。PyTorch 用了一个叫做版本计数器version counter的机制来检测这种情况。每个 Tensor 内部维护一个版本号每次原地操作都会让版本号递增。Autograd 引擎在反向传播时会检查当前 Tensor 的版本号是否和参与前向计算时一致不一致就直接报错。这种宁可报错也不给你算错的保守策略是 Autograd 安全性的底线。所以实操心法很简单前向计算图里凡是参与梯度计算的张量能不用原地操作就不用如果你想省显存做 inplace 的中间变量修改请先detach()切断梯度关联或者用克隆副本。尤其注意像nn.ReLU(inplaceTrue)这个经典参数它之所以能在很多模型里安全使用是因为 ReLU 的原地操作发生在激活值已经算出来之后而 PyTorch 的很多 Layer 内部实现已经考虑到了这种分支场景但你自己手写的那些、relu_()就要非常小心了。5. Activation 显存优化实战梯度检查点、手动释放与实操避坑5.1 用一个公式算清激活值显存你在为哪部分显存买单前面已经说了 Activation 是大头但很多人并不知道它的具体量级怎么估算。这里给一个简化版的公式Activation 显存 ≈ batch_size × 序列长度 × 隐藏维度 × 网络层数 × 每个元素字节数 × 系数以 GPT-2 规模的模型为例12 层、768 维隐藏层如果 batch size 是 8序列长度是 512那么单个 forward 过程中激活值的总体量约为$$8 \times 512 \times 768 \times 12 30,146,560 \times 4 \text{ bytes} \approx 120 \text{ MB}$$这还没算上 attention 的中间结果、FFN 扩展维度的中间张量。如果把多头注意力里的 Q、K、V 和注意力分数都算上激活值很容易再放大两三倍。对比一下模型参数本身——12 层 768 维参数量大约 1.17 亿FP16 存储约 234 MB。你会发现激活值和参数在显存里几乎是同一量级甚至更多。这就是为什么你在显存紧张时只缩小 batch size 比缩小模型本身更直接的原因。5.2 梯度检查点Gradient Checkpointing用重算换显存用时间换空间梯度检查点Gradient Checkpointing也叫 activation checkpointing是目前应对激活值显存爆炸最有效的手段之一在 HuggingFace Transformers 等库中已经变成标配选项。它的核心思路非常直白前向传播时不要保存所有层的激活值而是只保存少量检查点层的输出等到反向传播需要某个中间激活值时再临时从最近的检查点重新执行一遍前向把丢失的激活值重算出来。这个思路的本质是用计算换显存。重算激活值需要额外的 FLOPs但节省了显存占用。对于显存极度紧张、但 GPU 算力还有富余的场景这是性价比极高的折中方案。PyTorch 官方提供了现成的接口使用起来非常简单import torch.utils.checkpoint as checkpoint def forward(self, x): x checkpoint.checkpoint(self.block1, x) x checkpoint.checkpoint(self.block2, x) return x每个checkpoint调用都会创建一个前向区间。在这个区间内PyTorch 会故意不保存中间激活值只保存输入到该区间 Tensor 以及区间的输出。反向传播时会再次执行该区间的前向代码把激活值算回来再计算梯度。这里有几个实操中容易踩的坑。第一checkpoint函数要求被包裹的可调用对象是确定性的也就是说同样的输入必须产生同样的输出如果里面有随机 dropout梯度方向会出问题。解决办法是把随机种子传给函数或者使用checkpoint的preserve_rng_state参数来控制。想进一步省显存可以把它设为False。第二checkpoint对被包函数的参数数量有要求多参数函数请用lambda或functools.partial进行包装。第三checkpoint并不是无代价的它会让训练时间明显变长通常会增加 20% 到 30% 的时间开销在算力富余但显存吃紧的场景下收益极大。从显存账本的角度来看梯度检查点把激活值的峰值需求从 O(layer_count) 降到了 O(sqrt(layer_count))。因为理论上只要优化检查点的间隔就能把激活值的存储复杂度降到层数的平方根量级。这也是为什么那些几十层、上百层的大模型能够在一张消费级显卡上训练的关键原因之一。你完全可以把梯度检查点理解成按需重算的缓存淘汰策略跟操作系统的 swap 分页类似——内存不够就用 CPU 或者磁盘来换。5.3 手动释放与detach()的边界什么变量能删什么不能删除了梯度检查点日常训练中还有一些更轻量级的显存管理手段但用不好会适得其反。第一种是del手动删除中间变量配合torch.cuda.empty_cache()。这两个操作的作用经常被高估。del只是减少 Python 对象的引用计数真正的显存释放要等底层缓存池回收而torch.cuda.empty_cache()只是把显存缓存池里的空闲块返回给 CUDA 驱动并不会立刻减少显存占用。实际上 PyTorch 的显存分配器为了效率会缓存已释放的显存块下次分配直接复用所以频繁调用empty_cache()反而可能降低性能。真正有效的做法是在代码层面明确切断计算图。比如在训练循环里某个中间张量你已经在前向中算完了反向传播也用不到它了那就在下一次迭代开始前把它del掉或者用with torch.no_grad():包住推理过程。更系统的方法是利用detach()在不需要梯度的张量上切断计算图# 不对 z 保留梯度z 不再出现在计算图里 z y.detach()注意detach()返回的张量与原始张量共享数据内存但它不再参与梯度计算它的grad_fn为None。这在做特征提取、迁移学习、或者把某个模块的输出当成常数传给另一个模块时非常实用。还有一个经常被忽略的场景是验证集 / 测试集的前向。很多人写完model.eval()就完事但忘了前面那些张量仍然带着requires_gradTrue的计算图。推理时显存一点点被吃掉卡顿还不明显跑到后面突然 OOM。正确的做法是包上torch.no_grad()让推理阶段完全不构建计算图model.eval() with torch.no_grad(): for x, y in val_loader: pred model(x) ...这行代码能省下的显存可能比梯度检查点还多因为推理阶段连激活值都不需要保存了。5.4 实战组合拳梯度检查点、混合精度与梯度累积的一起调配在实际项目中显存优化很少只靠一种手段。我自己的经验是先用公式估算一下各部分占比再针对最肥的部分开刀。最经典的一套组合拳是混合精度训练 梯度检查点 梯度累积。混合精度AMP把前向和反向的矩阵计算改为 FP16直接减半 Activation 的字节数是最简单直接的显存红利。但要注意梯度累加时精度容易崩通常需要用 FP32 的 Master Weights 和 Loss Scaling 来控制。梯度检查点负责把 Activation 的存储量进一步压缩。梯度累积则解决 batch size 太大而放不下的问题它不减少单个 forward 的显存峰值但可以让有效的 batch size 变大。三者叠加效果往往能让一个本来 OOM 的模型从 24GB 降到 12GB 以下。我踩过的另一个坑是把 checkpoint 用在 batch 的第一层。那个被包进 checkpoint 的 module 如果输入是从 CPU 搬到 GPU 的大张量每次重算都要重新做一次 H2D 拷贝速度会慢得离谱。解决办法是把数据拷贝放在 checkpoint 外面或者用torch.utils.checkpoint的use_reentrant参数配合非重入模式来避免部分性能损失。另外PyTorch 2.x 里的torch.compile和 Dynamo 也集成了部分 activation checkpointing 的自动化能力但它的行为更多的是在图编译器层面做优化跟手写的torch.utils.checkpoint并不完全等价。我的建议是新手先从显式 checkpoint 入手把显存账本算清楚后再去探索自动化编译优化。最后补充一个排查 OOM 的通用心法不要急着改代码先搞清楚到底是谁在峰值时占着显存。可以写一小段测试脚本分别做只前向不反向和前向加反向两次实验用torch.cuda.max_memory_allocated()打点观察两次内存峰值的差异。差值越大的地方就是 Activation 越需要优化的地方。这套排查流程比盲调 batch size 有效得多。
返回列表