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_grad=True这个开关之后,它的内部结构瞬间复杂了一个维度。从 Autograd 的角度看,一个 Tensor 至少包含三样东西:
- data:真实的数值数据,也就是存储张量内容的内存空间;
- requires_grad:是否需要在该张量参与运算时追踪梯度;
- grad_fn:记录"这个张量是怎么算出来的"的函数对象。
新手最常见的误区,是把注意力全放在data上,觉得深度学习就是张量运算,而忽略了grad_fn才是 Autograd 的灵魂。grad_fn不是一个简单的属性标签,它是一个持有反向传播方法的函数节点。你执行y = x * w + b的时候,得到的y的grad_fn会指向一个乘法或加法的反向节点,这个节点记录了参与运算的输入 tensor 的引用。正因如此,前向过程结束后,你手上拿到的每一个中间张量,都自带一条"我从哪来"的完整记录。
可以做个简单实验来验证。定义两个叶子张量,做一次组合运算,然后打印中间结果的grad_fn:
import torch x = torch.randn(3, 3, requires_grad=True) w = torch.randn(3, 3, requires_grad=True) 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_grad=True且是叶子节点,所以它们的.grad会被填充并用于参数更新。 - 叶子节点在创建时如果
requires_grad=True,它的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_grad=True) 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_grad=True) 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_grad=True) 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(inplace=True)这个经典参数,它之所以能在很多模型里安全使用,是因为 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_grad=True的计算图。推理时显存一点点被吃掉,卡顿还不明显,跑到后面突然 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 有效得多。