这两年不管是跑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_graph=True)。这种“用完即弃”的特性,让PyTorch在动态控制流场景下拥有压倒性优势——模型的输入长度、循环次数、条件分支每次前向都可能不同,动态图天然适配这种需求,不需要额外的trace机制来适配。
1.2 从用户代码到图的诞生:一次前向传播全记录
我习惯在讲计算图时用一段最简单的代码来演示,因为只有亲手打印出这些东西,你才能真正感受到“代码即图”的含义。看下面这个例子:
import torch # 叶子节点:用户创建且 requires_grad=True x = torch.randn(4, 8, requires_grad=True) w = torch.randn(8, 12, requires_grad=True) b = torch.randn(12, requires_grad=True) # 非叶子节点:由算子计算得到 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^T
PyTorch的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_fn | None | 有(记录操作) | 非叶子节点必有反向函数 |
is_leaf | True | False | 结构层面的标记 |
.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_size=8、seq_len=512、hidden_size=1024、FFN中间层是4倍隐藏维度、模型12层:
- 单层Attention的Q/K/V输出:3 × 8×512×1024×4字节 ≈ 50MB
- Attention score矩阵:8×8(heads)×512×512×4字节 ≈ 67MB
- FFN第一层输出:8×512×4096×4字节 ≈ 67MB
- FFN第二层输出:8×512×1024×4字节 ≈ 16MB
- 单层activation合计约200MB,12层就是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, device='cuda') print(f"initial: {mem_mb():.1f} MB") loss = model(x).square().mean() print(f"after forward: {mem_mb():.1f} MB") loss.backward() print(f"after backward: {mem_mb():.1f} MB") optimizer.step() print(f"after step: {mem_mb():.1f} MB")通常你会看到这样的现象:前向结束后显存跳升一大截,反向结束后显存又降回来一部分,但不会完全回到前向之前的状态。原因是模型参数的梯度、优化器状态这些是“持久化”的,会一直占用显存直到训练结束。而中间activation是“临时”的,反向用完之后就被释放了。
这个“前向涨、反向跌”的过程,就是计算的显存生命周期最直接的体现。
4. Activation显存优化:把省显存这件事做到极致
4.1 梯度检查点:拿计算换显存
清楚了activation是显存大头,优化策略就有的放矢了。最常用的手段是梯度检查点(gradient checkpointing),核心思想很简单:前向传播时不保存中间activation,只保存这一层的输入;反向传播需要中间结果时,临时重新执行一次前向计算出来。这就是典型的“用时间换空间”。
PyTorch提供了非常方便的API:
from 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_grad=False,反向梯度不会再流回原来的图。典型误用场景是:你想把某个中间特征存下来做可视化或写日志,如果直接保存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_graph=True。这个参数会让计算图在反向后保留,正常情况下完全不需要。 - 检查验证阶段是否忘记加
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_grad=True的张量执行了类似x += 1、x.zero_()、relu_()等inplace操作。排查思路:
- 在报错栈中找到触发inplace操作的代码行,重点检查
+=、*=、zero_、fill_这类写法。 - 把inplace改成out-of-place版本,比如
x = x + 1而不是x += 1。 - 如果确实需要inplace,确保操作发生在
no_grad上下文里,并且不会影响反向计算需要的前向值。比如对不需要梯度的输入做inplace是安全的,但对应requires_grad=True的叶子节点操作就要高度警惕。
这个报错有时候非常难排查,因为触发点可能在模型内部、损失函数里,或者是在自定义层里。我的建议是一旦出现这类报错,优先检查最近改动的代码块,尤其是那些为了“省显存”而改成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机制,初看是理论问题,实际用起来全是显存问题。把这套生命周期弄明白,你的训练稳定性至少能上一个台阶。