☰
PyTorch内部机制深度解析:从Tensor、Autograd到算子执行
2026/10/1 17:00:46 网站建设 项目流程

PyTorch 用久了,总会有那么一个时刻,你盯着报错信息发呆:明明张量形状对得上,梯度却传不回去;或者loss.backward()跑完,某个中间变量的.grad是None。这时候翻文档往往只能查到 API 签名,真正想搞明白“它内部到底怎么跑的”,还是得把 Autograd、Tensor、Storage 这几层拆开看。这篇内容就是围绕 PyTorch 的内部机制做一次系统梳理,从张量和存储的分离设计,到动态计算图的构建,再到反向传播时梯度是怎么一层层算出来的,最后落到算子层面的执行流程。适合已经能跑通训练脚本、但想进一步理解框架行为、排查梯度异常、做自定义算子或性能优化的读者。下面这些内容,一部分来自源码阅读,一部分来自实际调试中踩过的坑,尽量说人话,把“为什么这么设计”讲清楚。

1. 从一次梯度为 None 的排查说起

1.1 问题的表象与第一反应

之前有个朋友拿了一段代码来问,说模型训练不收敛,检查了半天发现某个中间层的权重梯度是None。他的第一反应是“是不是这个层没参与计算”,于是打印了前向输出,发现输出是正常的,数值也在变。这就很反直觉:既然前向参与了,为什么反向没有梯度?

我让他把计算图打印出来看,结果发现那个权重虽然参与了前向,但它的计算路径上有一个detach()操作。detach()会把一个张量从计算图中摘出来,后续所有基于它的运算都不会再记录梯度。前向数值照样算,因为数值计算和梯度记录是两条线。这就是很多人第一次接触 Autograd 时容易混淆的点:前向传播和梯度追踪不是一回事。

这个案例其实暴露了一个核心问题:如果不理解 PyTorch 内部是怎么组织张量、怎么构建计算图、怎么在反向时遍历图,遇到这类问题就只能靠猜。而一旦把内部机制理清楚,这类问题基本是看一眼就能定位。

1.2 为什么值得花时间理解内部机制

有人会说,框架封装好了,能用就行,何必关心内部。这话在大多数场景下没错,但有几类情况绕不开:

  • 调试梯度异常:梯度为 None、梯度爆炸、梯度被意外截断,这些问题的根因往往在计算图的构建阶段,而不是数值本身。
  • 自定义算子:写torch.autograd.Function的时候,必须手动实现forward和backward,不理解 Autograd 的调度逻辑根本写不对。
  • 性能优化:知道 Tensor 和 Storage 的关系,才能理解为什么view比reshape快、为什么原地操作有时会报错。
  • 模型部署:导出 ONNX 或者做图优化时,计算图的结构直接决定了能不能导出、导出后对不对。

所以这篇内容不是纯理论,而是围绕“能解决实际问题”来组织的。下面从最底层的存储结构开始,一层层往上拆。

2. Tensor 与 Storage 的分离设计

2.1 为什么张量不直接持有数据

刚接触 PyTorch 的人通常会以为一个 Tensor 就是一块内存加上形状信息。实际上 PyTorch 把这两者拆开了:Tensor 负责描述“怎么看待数据”,Storage 负责“数据存在哪”。一个 Tensor 包含形状(size)、步长(stride)、偏移(offset)、数据类型(dtype)等信息,而真正的数值存在一个连续的 Storage 里。

这么设计的好处很直接:多个 Tensor 可以共享同一块 Storage,只是用不同的形状和步长去“解读”它。最典型的就是view和transpose。transpose不会复制数据,它只是把步长换了一下,底层 Storage 完全没动。你可以用下面这段代码验证:

import torch a = torch.arange(12).reshape(3, 4) b = a.transpose(0, 1) print(a.storage().data_ptr() == b.storage().data_ptr()) # True print(a.stride(), b.stride()) # (4, 1) (1, 4)

两个张量的 Storage 指针完全一样,说明它们共享内存。区别只在 stride:a的行步长是 4、列步长是 1,b反过来。这就是为什么transpose几乎不耗时,而contiguous()会触发一次真正的内存拷贝。

2.2 步长、偏移与视图的边界

理解了 stride,很多“反直觉”的行为就说得通了。比如a[1:]这种切片,它不会复制数据,而是返回一个偏移了若干字节、形状变小的视图。偏移量(storage_offset)记录的就是这个视图从 Storage 的哪个位置开始。

这里有个容易踩的坑:视图操作和原地操作混用时,可能改到不该改的数据。举个例子:

a = torch.arange(12).reshape(3, 4) b = a[0] # b 是 a 第一行的视图 b[0] = 999 # 原地修改 b print(a[0, 0]) # 999,a 也被改了

因为b和a共享 Storage,改b就是改a。这在写数据处理管道时特别容易出问题,尤其是把切片结果传给别的函数做原地操作。我的习惯是,只要一个张量会被原地修改,就先.clone()一份,虽然多一次拷贝,但能避免很多隐蔽的 bug。

2.3 view 与 reshape 的本质区别

view和reshape看起来功能一样,都是改形状,但内部行为不同。view要求张量在内存里是连续的(或者至少满足特定的 stride 条件),因为它只是重新解释 stride,不搬数据。如果张量不连续,view会直接报错。reshape则更宽容:能 view 就 view,不能 view 就先拷贝成连续的再 view。

a = torch.arange(12).reshape(3, 4) b = a.transpose(0, 1) # b.view(-1) # 报错,b 不连续 c = b.reshape(-1) # 可以,内部先 contiguous 再 view

所以性能敏感的代码里,如果确定张量连续,用view更明确;如果不确定,用reshape更安全,但要知道它可能偷偷做一次拷贝。这个区别在做大张量操作时影响很明显,一次不必要的拷贝可能就是几十毫秒。

3. Autograd 的动态图构建过程

3.1 计算图不是预先定义的

PyTorch 和早期的一些框架最大的区别在于:它的计算图是动态构建的,也就是在每次前向传播时现场生成。你写一行运算,它就往图里加一个节点。这也是为什么 PyTorch 里可以用 Python 的if、for控制流,因为图是跟着代码执行走的。

每个需要梯度的张量都有一个grad_fn属性,指向创建它的那个函数节点。叶子节点(比如模型参数)的grad_fn是None,非叶子节点的grad_fn记录了它是怎么算出来的。反向传播时,从 loss 出发,沿着grad_fn链一路往回走,这就是链式法则的工程实现。

x = torch.tensor([2.0], requires_grad=True) y = x ** 2 z = y * 3 print(x.grad_fn) # None,叶子节点 print(y.grad_fn) # <PowBackward0> print(z.grad_fn) # <MulBackward0>

从z的grad_fn出发,能找到y,再找到x,这条链就是计算图。理解这一点,就能明白为什么detach()有效:它把链断开了,后续节点不再指向原来的图。

3.2 requires_grad 的传播规则

一个张量是否需要梯度,由requires_grad决定。这个属性会沿着计算传播:只要有一个输入需要梯度,输出通常就需要梯度。但有几个例外需要记住:

  • 整数类型的张量不能要求梯度,这是硬性限制。
  • 如果所有输入都不需要梯度,输出也不需要。
  • 在torch.no_grad()上下文里,即使输入需要梯度,输出也不会记录梯度。

torch.no_grad()在推理阶段非常常用,它不只是省显存,更重要的是避免构建无用的计算图。我见过有人在验证循环里忘了加no_grad,结果显存一路涨,最后 OOM。原因就是每次验证都在建图,图越积越多。

3.3 叶子节点与非叶子节点的梯度

默认情况下,只有叶子节点的梯度会被保留在.grad里,非叶子节点的梯度算完就释放了。这是为了省内存。如果你想看中间某个张量的梯度,得手动调用retain_grad():

x = torch.tensor([2.0], requires_grad=True) y = x ** 2 y.retain_grad() z = y * 3 z.backward() print(y.grad) # tensor([3.]) print(x.grad) # tensor([12.])

这个机制在调试时特别有用。很多人调试梯度问题时,直接打印中间变量的.grad发现是None,就以为梯度没传过去,其实只是被释放了。加上retain_grad()再看,往往就正常了。

4. 反向传播的调度与梯度累加

4.1 backward 到底做了什么

调用loss.backward()时,PyTorch 做的是:从 loss 这个节点出发,按拓扑逆序遍历计算图,对每个节点调用它对应的反向函数,把上游传来的梯度乘以本地的雅可比矩阵,再传给下游。这个过程是自动的,但有几个细节值得注意。

首先是拓扑排序。计算图可能有分支和合并,必须保证一个节点的所有下游梯度都到齐了才能算它自己的梯度。PyTorch 用拓扑排序保证这个顺序。如果图里有环(正常前向不会产生),反向就会出问题。

其次是梯度累加。如果一个张量被多条路径用到,它的梯度是各路径梯度之和。这就是为什么backward()默认是累加而不是覆盖。很多人训练时忘了zero_grad(),梯度就一直累加,导致更新步长越来越大,loss 直接飞掉。

optimizer.zero_grad() # 清空上一轮梯度 loss.backward() # 累加本轮梯度 optimizer.step() # 更新参数

这三行的顺序不能乱。zero_grad必须在backward之前,step必须在backward之后。

4.2 梯度累加的实际用途

梯度累加虽然容易踩坑,但它本身是个有用的特性。当显存不够、没法开大 batch 时,可以用“小 batch 多次前向 + 梯度累加”来模拟大 batch:

for i, (data, target) in enumerate(loader): output = model(data) loss = criterion(output, target) / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()

这里把 loss 除以累加步数,是为了让梯度的量级和真实大 batch 一致。如果不除,累加后的梯度会偏大,相当于变相提高了学习率。这个技巧在显存受限时非常实用,但要注意 BatchNorm 这类依赖 batch 统计的层,小 batch 下的统计量和大 batch 不一样,效果可能有差异。

4.3 高阶梯度与 create_graph

PyTorch 还支持高阶梯度,也就是对梯度再求梯度。这需要backward()时传create_graph=True,让反向过程本身也被记录成图:

x = torch.tensor([2.0], requires_grad=True) y = x ** 3 dy = torch.autograd.grad(y, x, create_graph=True)[0] d2y = torch.autograd.grad(dy, x)[0] print(d2y) # 12.0,即 6x

高阶梯度在实现某些正则化、元学习或者物理约束的损失时会用到。但它的开销比一阶大不少,因为反向图也要建、也要占显存。不是必需就别开。

5. 算子层面的执行流程

5.1 一个算子从调用到执行经历了什么

当你在 Python 里写torch.add(a, b)或者a + b时,背后经历了一条不短的链路。简单说:Python 层调用进入 C++ 的 dispatcher,dispatcher 根据设备类型(CPU/CUDA)、数据类型、是否需要梯度等信息,选择合适的 kernel 去执行。这个分发机制叫dispatch,是 PyTorch 支持多后端的关键。

以加法为例,如果两个张量都在 CUDA 上,dispatcher 会路由到 CUDA 的加法 kernel;如果在 CPU 上,路由到 CPU kernel;如果涉及自动微分,还会在计算图里注册一个AddBackward节点。这一整套流程对用户是透明的,但理解它有助于排查“为什么这个算子在 GPU 上没生效”这类问题。

5.2 算子融合与性能

大量小算子的连续调用是性能杀手,因为每个算子都有启动开销,GPU 上尤其明显。PyTorch 2.0 引入的torch.compile就是干这个的:把一段计算图编译融合,减少 kernel 启动次数,同时做算子级别的优化。

即使不用torch.compile,手动减少算子数量也有收益。比如把a * b + c写成torch.addcmul(c, a, b),虽然语义一样,但后者是一个融合算子,少一次中间结果的读写。在元素级操作密集的模型里,这种优化累积起来很可观。

5.3 自定义算子的两种方式

需要写自定义算子时,有两条路:

  • torch.autograd.Function:适合需要自定义前向和反向逻辑的场景。你要手动实现forward和backward,backward里返回对每个输入的梯度。
  • 扩展 C++/CUDA:适合性能敏感、需要底层优化的场景。通过torch.utils.cpp_extension编译自定义 kernel。

用Function写自定义算子时,最容易出错的是backward的返回值个数和顺序必须和forward的输入一一对应,而且只有requires_grad=True的输入才需要返回梯度,其他的返回None。这个规则不遵守,反向就会报错或者静默出错。

class MyReLU(torch.autograd.Function): @staticmethod def forward(ctx, x): ctx.save_for_backward(x) return x.clamp(min=0) @staticmethod def backward(ctx, grad_output): x, = ctx.saved_tensors return grad_output * (x > 0).float()

ctx.save_for_backward用来保存反向需要的前向中间结果,比直接存在ctx属性上更安全,因为它能正确处理内存和版本管理。

6. 几个实际调试中的经验

6.1 定位梯度问题的通用思路

遇到梯度异常,我一般按这个顺序排查:先确认requires_grad有没有被意外关掉,再检查计算图里有没有detach或no_grad断链,然后看是不是非叶子节点的梯度被释放了(加retain_grad),最后才怀疑数值问题。大部分“梯度为 None”的问题都出在前三步,真正数值层面的问题反而少。

6.2 原地操作的版本检查

PyTorch 对原地操作有版本检查机制。如果一个张量在被用于计算后又被原地修改,反向传播时可能用到错误的数据,PyTorch 会直接报错提示版本不匹配。这个报错看着吓人,其实是在保护你。解决办法通常是避免原地操作,或者调整操作顺序。

6.3 显存与计算图的释放

计算图在backward()之后默认会被释放,这也是为什么反向只能调用一次。如果想多次反向(比如某些对抗训练场景),需要传retain_graph=True。但要注意,保留图会一直占显存,用完记得手动释放。我见过有人为了图省事到处加retain_graph=True,结果显存泄漏,训练跑一半就崩。

理解 PyTorch 内部机制这件事,投入产出比其实挺高的。花几个小时把 Tensor、Storage、Autograd、算子这几层的关系理清楚,后面遇到问题时定位速度会快很多,写自定义算子和做性能优化时也更有底气。我自己的习惯是,每学一个新框架,都先把它最核心的那两三个抽象搞明白,剩下的 API 都是在这上面长出来的。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询