
这两年不管是跑CV还是NLP模型我经常被问到同一个问题训练刚开始显存直接飙满等loss.backward()跑完显存又刷刷往下降这是为什么很多人第一反应是模型参数太多但其实模型参数只占显存的一小部分真正吃显存的是前向过程留下的中间结果——也就是activation。要搞清楚显存在哪儿跑哪去了为什么反向之后才释放就得把PyTorch的计算图机制和Autograd的底层工作方式从头捋一遍。这篇笔记我就围绕动态DAG、反向求导机制和activation显存优化这条主线把我实际调试中总结的经验一次说清楚。PyTorch的计算图和Autograd不是两套独立的东西它们本质上是同一套机制的一体两面前向计算时构建计算图反向传播时沿着计算图求梯度。理解这一点之后你再看显存问题视角会完全不一样——显存不是在某个瞬间“突然爆掉”的而是随着计算图的构建、保留、释放经历了一个完整的生命周期。把这张图的生命周期吃透显存优化就不再是盲人摸象了。如果你正准备入门PyTorch或者已经在训练模型但经常被OOM和梯度问题折磨这篇内容应该能帮你省下不少排查时间。我不会堆概念尽量用代码和实际经验来讲每一步都能直接动手验证。1. 先看整体设计计算图与Autograd是一套什么机制1.1 动态DAG的核心设计思路PyTorch的autograd本质上就是一个自动差分引擎它在用户执行前向计算的过程中把所有参与计算且需要梯度的张量及其运算关系记录下来构成一个有向无环图。这个DAG的节点是Tensor边是运算逻辑数据流方向就是前向传播的方向。这里的关键词是“有向无环”——数据只能从输入流向输出不能回环这是反向传播能够稳定执行的前提。torch.Tensor内部都有一个grad_fn字段每个算子比如add、mm、relu在计算出输出张量的同时会创建一个对应的反向函数记录进图里。换句话说PyTorch是“边算边记”的代码执行到哪里图就构建到哪里。这个设计与传统静态图框架有本质区别。静态图的做法是先定义图结构再喂数据执行图一旦建好就不能改了。动态图则是每次前向传播时重新构建一张全新的图反向传播结束后这张图会被释放除非显式设置retain_graphTrue。这种“用完即弃”的特性让PyTorch在动态控制流场景下拥有压倒性优势——模型的输入长度、循环次数、条件分支每次前向都可能不同动态图天然适配这种需求不需要额外的trace机制来适配。1.2 从用户代码到图的诞生一次前向传播全记录我习惯在讲计算图时用一段最简单的代码来演示因为只有亲手打印出这些东西你才能真正感受到“代码即图”的含义。看下面这个例子import torch # 叶子节点用户创建且 requires_gradTrue x torch.randn(4, 8, requires_gradTrue) w torch.randn(8, 12, requires_gradTrue) b torch.randn(12, requires_gradTrue) # 非叶子节点由算子计算得到 y x w b loss y.square().mean() print(loss.grad_fn:, loss.grad_fn) # MeanBackward0 object at ... print(y.grad_fn:, y.grad_fn) # AddBackward0 object at ... print(x.is_leaf:, x.is_leaf) # True print(y.is_leaf:, y.is_leaf) # False print(loss.requires_grad:, loss.requires_grad) # True执行完这几行计算图就已经构建完成了。图中的结构大概是这样的loss节点依赖mean运算mean依赖pow运算pow依赖add运算add依赖mm运算。每一个grad_fn都对应一个反向函数它们串联起来的顺序就是将来反向传播要走的路径。我在实际项目中调试梯度问题时经常用这种打印grad_fn的方式快速定位计算图的形态。比如当你怀疑某个变量没有梯度时先看它是不是叶子节点再看它的grad_fn是什么通常很快就能找到问题。1.3 为什么是动态图与静态图的取舍很多人以为动态图是PyTorch的“唯一正确选择”其实不对。动态图的好处是灵活、易调试代码怎么写图就是什么样Python层面的调试工具全部通用。但代价是每次前向都要重新建图并且图的信息是局部的——框架很难做跨算子的全局优化比如算子融合、内存规划这类静态图才能施展的优化手段。静态图的好处恰恰是动态图缺的一次建图、多次执行执行引擎可以对整张图做全局分析和内存规划。这也是为什么PyTorch后来要推出torch.compile和TorchScript——它本质上是想在不牺牲动态灵活性的前提下补上静态优化的能力。理解这个背景对后续显存优化也有帮助。因为torch.compile在做内存规划时确实能比纯动态图更高效地复用显存。但回归到日常训练大部分时候我们还是在纯动态图模式下运行所以理解autograd的显存行为依然是基本功。2. 反向求导的底层逻辑梯度是如何沿着DAG跑回去的2.1 链式法则在DAG上的具体表达反向传播不是“玄学”本质就是微积分的链式法则。如果把DAG看作一条流水线前向是从输入流向输出反向就是梯度从输出流回输入每经过一个节点就乘以这个节点的局部梯度。我用刚才的例子具体算一遍。假设我们有一个非常简单的复合函数z x w # 矩阵乘法 t z b # 加法 u t ** 2 # 平方 loss u.mean() # 均值反向传播时框架要依次计算d(loss)/d(u)mean的局部梯度每个元素都是1/Nd(loss)/d(t)乘以du/dt 2td(loss)/d(z)加法的局部梯度是1直接透传d(loss)/d(x)矩阵乘法的局部梯度需要乘以w^Td(loss)/d(w)同理需要乘以x^TPyTorch的autograd引擎做的事情就是对这个DAG做一次拓扑排序然后从输出节点开始逆向调用每个grad_fn的backward()方法把梯度一步一步“分发”回各个输入。这就是整个反向求导机制的机械实现——每一步都是明确、可验证的。2.2 Autograd内部的传递过程与节点状态loss.backward()执行时实际发生了这几件事从当前节点开始拿到初始种子梯度。对于标量loss默认梯度是1.0。对图进行拓扑排序确定反向遍历的顺序。注意这个顺序必须严格保证每个节点的依赖都已经计算出梯度后才能轮到它自己。对每个节点的grad_fn调用backward()计算出对各个输入的梯度贡献。把梯度累加到对应Tensor的.grad属性中。这里要特别强调“累加”两个字——autograd不会替你做梯度清零所以训练循环里必须在每次optimizer.step()之前手动optimizer.zero_grad()否则梯度会不断叠加。反向传播过程中如果某个中间节点的依赖已经全部完成而且这个节点没有被要求保留autograd就会把它的缓存释放掉。这个过程解释了很多人遇到过的现象backward()结束之后中间层的显存占用明显下降。那不是错觉而是autograd主动清理了中间激活值。2.3 grad_fn、is_leaf与requires_grad三个概念的辨析这三个属性是PyTorch新手最容易搞混的地方我直接列一个表格对比清楚属性叶子节点非叶子节点说明requires_grad可手动设置自动继承输入决定是否进入自动微分grad_fnNone有记录操作非叶子节点必有反向函数is_leafTrueFalse结构层面的标记.grad有默认为空非叶子节点默认不保留梯度很多项目中开发者会在非叶子张量上调用.grad发现是None以为梯度丢了。其实不是这是autograd的默认行为——为了省显存和计算中间变量的梯度不会保留除非你显式调用.retain_grad()。在调试梯度流的时候这是一个很实用的辅助手段y (x * 2).sum() y.retain_grad() # 强制保留非叶子节点的梯度3. 显存生命周期前向存、反向还的完整过程3.1 PyTorch显存管理器的工作方式讲到显存生命周期必须先把PyTorch的底层显存管理机制讲清楚。PyTorch并不是“用多少显存就实时向GPU申请多少”而是通过一个CachingAllocator向CUDA驱动申请一大块显存预留下来Tensor销毁后显存块并不是直接还给驱动而是回到缓存池复用。这个设计的好处是显存分配速度极快——从缓存池里拿一块现成的比向驱动申请快几个数量级。但副作用也明显你在任务管理器里看到的GPU显存占用不会因为某个Tensor被释放就立刻降下来。只有当缓存池里的空闲块长时间没有被再次使用或者你手动调用torch.cuda.empty_cache()这些显存才会真正归还给CUDA。这个机制还解释了另一个现象频繁创建大小不一的Tensor会导致显存碎片化。缓存池里会出现大量大小不匹配的空闲块新的大Tensor申请不到连续显存于是触发新的驱动申请显存峰值越拉越高。3.2 activation是显存大头算一笔账我用Transformer为例实际估算一下训模型时显存到底花在哪。假设batch_size8、seq_len512、hidden_size1024、FFN中间层是4倍隐藏维度、模型12层单层Attention的Q/K/V输出3 × 8×512×1024×4字节 ≈ 50MBAttention score矩阵8×8(heads)×512×512×4字节 ≈ 67MBFFN第一层输出8×512×4096×4字节 ≈ 67MBFFN第二层输出8×512×1024×4字节 ≈ 16MB单层activation合计约200MB12层就是2.4GB以上作为对比这12层Transformer的参数总量可能只有几百MB。所以在训练阶段activation才是显存开销的真正大头这个问题在推理阶段完全感知不到因为推理时no_grad模式下不需要保存梯度相关的中间结果所以很多人在部署时觉得模型很“轻”一训练就露馅。3.3 反向传播前后的显存变化实测理论讲再多不如动手跑一段代码实测。我在调试显存问题时常用下面这段代码来观察显存变化import torch def mem_mb(): return torch.cuda.memory_allocated() / 1024 ** 2 model torch.nn.Linear(1024, 1024).cuda() optimizer torch.optim.Adam(model.parameters()) x torch.randn(64, 1024, devicecuda) print(finitial: {mem_mb():.1f} MB) loss model(x).square().mean() print(fafter forward: {mem_mb():.1f} MB) loss.backward() print(fafter backward: {mem_mb():.1f} MB) optimizer.step() print(fafter step: {mem_mb():.1f} MB)通常你会看到这样的现象前向结束后显存跳升一大截反向结束后显存又降回来一部分但不会完全回到前向之前的状态。原因是模型参数的梯度、优化器状态这些是“持久化”的会一直占用显存直到训练结束。而中间activation是“临时”的反向用完之后就被释放了。这个“前向涨、反向跌”的过程就是计算的显存生命周期最直接的体现。4. Activation显存优化把省显存这件事做到极致4.1 梯度检查点拿计算换显存清楚了activation是显存大头优化策略就有的放矢了。最常用的手段是梯度检查点gradient checkpointing核心思想很简单前向传播时不保存中间activation只保存这一层的输入反向传播需要中间结果时临时重新执行一次前向计算出来。这就是典型的“用时间换空间”。PyTorch提供了非常方便的APIfrom torch.utils.checkpoint import checkpoint def transformer_layer(x, attn, ffn): # 把每一个子模块包进 checkpoint 里 x checkpoint(attn, x) x checkpoint(ffn, x) return x这种方式在Transformer大类模型上效果极其显著通常能砍掉一半以上的activation显存在显存不足时是救命稻草。代价是训练时间会增加30%~50%因为反向时要重复计算被“检查点”包住的那部分前向逻辑。我的一个经验是checkpoint的分段粒度要适中。如果你把每一个微小算子都单独包进checkpoint重算开销会大得离谱收益反而下降。合理做法是让每个checkpoint包住一个完整的模块——比如整个Attention块或整个FFN块。另外如果模型中没有dropout这类随机操作可以把checkpoint的preserve_rng_state参数设成False能省下一点点RNG状态保存和恢复的开销。4.2 no_grad、detach、del三种切断计算图的方式除了梯度检查点实际项目里我更常用的其实是“切断计算图”的思路。这有三种手段适用场景完全不同。第一种是torch.no_grad()。在推理、评估或者只需要提取特征不更新梯度的阶段用no_grad上下文包裹代码PyTorch就直接不构建计算图了。这是零成本省显存的方式也是很多人在验证集上忘记加no_grad导致显存爆炸的常见原因。第二种是.detach()。它的作用是从计算图中把一个张量“摘出来”返回的新张量requires_gradFalse反向梯度不会再流回原来的图。典型误用场景是你想把某个中间特征存下来做可视化或写日志如果直接保存y而不是y.detach()那么这个张量会一直持有整张计算图的引用显存怎么都降不下去。正确做法是# 错误示范保存带梯度的张量整张图被长周期持有 all_losses.append(loss) # 正确示范只保存数值切断图引用 all_losses.append(loss.item())第三种是del。Python层面删除变量引用让Tensor的引用计数归零底层显存就能及时释放回缓存池。注意它只对“释放引用”有效如果张量本身还在计算图里或者被其他容器引用着del也拿它没办法。所以在必要的时候del配合torch.cuda.empty_cache()使用效果更好。4.3 混合精度、梯度累积与inplace操作的实际取舍这三个手段在显存优化中各有位置但都伴随代价要根据场景谨慎选择。混合精度AMP是目前性价比最高的方案。把模型权重和activation以FP16存储显存直接减半同时计算速度还有提升。但要注意两点一是梯度下溢问题FP16能表示的数值范围有限反向传播时梯度可能小到变成0所以需要用torch.cuda.amp.GradScaler做loss scaling二是某些对精度极其敏感的操作比如LayerNorm、Softmax建议保持在FP32下计算避免数值不稳定。PyTorch的autocast会自动处理大部分情况但你要知道背后发生了什么否则遇到精度问题时无从下手。梯度累积gradient accumulation的思路是把大batch拆成几个小batch每步都做前向和反向但是不立即更新参数而是在累积到一定梯度数量后再执行optimizer.step()。从显存角度看每个小batch的activation在反向后就被释放所以显存峰值跟单个小batch一致不需要一次性预留大batch的activation。这个方案几乎不损失精度唯一的代价是训练时间略增而且需要小心BatchNorm这类依赖统计量的层在梯度累积模式下行为可能不太一样。inplace操作是三者中最需要谨慎的。relu_、add_这类inplace操作把输出写回输入张量从原理上确实能省一份中间显存。但风险在于它可能覆盖掉反向计算时需要的前向输入值导致梯度计算错误。PyTorch虽然对inplace有检测但报错信息往往不够直观。我的原则很简单不明确知道某个inplace不会影响反向传播链的前提下一律用out-of-place版本。5. 实战中的常见问题与避坑记录5.1 模型明明不大却反复OOM这是被问得最多的问题。遇到OOM先别急着怀疑模型太大按照我的排查顺序走一遍检查训练循环里是不是把带梯度的Tensor存进了容器。比如用List累积loss时用loss而不是loss.item()整张计算图都会被长周期持有。检查backward()是否误用了retain_graphTrue。这个参数会让计算图在反向后保留正常情况下完全不需要。检查验证阶段是否忘记加no_grad()。验证时不需要梯度但不加no_grad会构建一套额外计算图白占显存。检查DataLoader的num_workers是否设置过高。加载数据的进程本身也可能占用显存尤其是在GPU上做数据增强时。检查是否在循环里不断创建新的Tensor而没有释放旧引用。有时候不是单一大Tensor的问题而是小Tensor累积造成的碎片化。我印象很深的一次事故训练时为了画loss曲线把每个step的loss张量都append到一个List里结果显存只涨不降整个训练在第几百个step之后必挂。改成loss.item()存储后问题瞬间消失。这种低级错误其实相当常见。5.2 inplace操作导致“gradient computation has been modified”这个报错信息非常经典one of the variables needed for gradient computation has been modified by an inplace operation。意思是反向传播需要的前向输入值已经被某个inplace操作覆盖了autograd无法正确计算梯度。最常见的原因是对requires_gradTrue的张量执行了类似x 1、x.zero_()、relu_()等inplace操作。排查思路在报错栈中找到触发inplace操作的代码行重点检查、*、zero_、fill_这类写法。把inplace改成out-of-place版本比如x x 1而不是x 1。如果确实需要inplace确保操作发生在no_grad上下文里并且不会影响反向计算需要的前向值。比如对不需要梯度的输入做inplace是安全的但对应requires_gradTrue的叶子节点操作就要高度警惕。这个报错有时候非常难排查因为触发点可能在模型内部、损失函数里或者是在自定义层里。我的建议是一旦出现这类报错优先检查最近改动的代码块尤其是那些为了“省显存”而改成inplace的地方。5.3 梯度检查点开启后训练明显变慢这是完全正常的现象但如果慢得离谱多半是检查点粒度设置不当。如果你把每个小算子都包进checkpoint重算开销会指数级上升因为每一层都变成了“前向重算反向再算”的多倍开销训练时间可能翻倍甚至更多。合理做法是对显存占比最高的模块启用检查点。以Transformer为例通常只把FFN块和Attention块分别包起来就够了不需要深入到每个矩阵乘法层面。另一个优化点是preserve_rng_state参数。如果模型中没有dropout、RandomCrop这类随机操作设成False既可以减少RNG状态的保存与恢复开销也能省一点点显存。如果你发现checkpoint带来的收益不明显可以先关掉它用torch.cuda.memory_summary()看看每个阶段的显存占用分布再决定到底该不该上checkpoint、包住哪部分。盲目开所有优化手段有时候反而适得其反。5.4 显存碎片化与缓存清理技巧训练过程中显存占用没有明显下降但降低batch size或重启训练却总是触发OOM大概率是显存碎片化在作怪。CachingAllocator的缓存池里积累了太多大小不一的空闲块新请求找不到连续空间时就得向驱动申请新的显存块于是峰值越来越高。一个直接的缓解手段是定期调用torch.cuda.empty_cache()把缓存池中的空闲块归还给CUDA驱动。要注意empty_cache()只回收“空闲”的块正在使用的Tensor不受影响。如果你用了checkpoint它注册的前向重算备份可能正在占用显存调用empty_cache()并不会帮你释放正在被图引用的部分。更稳妥的方案是减少显存分配的“碎片化来源”尽量固定Tensor的形状避免在循环里频繁创建尺寸变化很大的Tensor长度不一的序列用padding对齐减少不必要的梯度累积中间态。这些做法比事后清理缓存更根本。另外一个比较实用的小技巧是直接用torch.cuda.memory_summary()打印完整的显存分配报告能清楚看到哪一层、哪一步占了多大显存排查效率会高很多。。。。继续写结尾部分踩过的坑多了之后我越来越觉得显存问题本质上就是计算图生命周期的问题。很多人一上来就想开AMP、开checkpoint但我建议你先搞清楚显存到底花在哪。用memory_summary()看一眼再看看是不是有带梯度的Tensor被存进了某个List是不是验证阶段忘了no_grad。很多时候不是模型太大而是计算图的引用没有及时释放。最后再分享一个小技巧在训练循环里定期打印torch.cuda.memory_summary()它会显示每个阶段的显存分配和缓存池状态是我排查显存问题时用的最多的工具。PyTorch的计算图和autograd机制初看是理论问题实际用起来全是显存问题。把这套生命周期弄明白你的训练稳定性至少能上一个台阶。