尧图精选

PyTorch内部机制深度解析:从Tensor、Autograd到算子执行

🕒 发布时间:2026/10/1 17:00:48 📁 来源:尧图网络
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 指针完全一样说明它们共享内存。区别只在 stridea的行步长是 4、列步长是 1b反过来。这就是为什么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]) # 999a 也被改了因为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_gradTrue) 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_gradTrue) 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 多次前向 梯度累加”来模拟大 batchfor 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_graphPyTorch 还支持高阶梯度也就是对梯度再求梯度。这需要backward()时传create_graphTrue让反向过程本身也被记录成图x torch.tensor([2.0], requires_gradTrue) y x ** 3 dy torch.autograd.grad(y, x, create_graphTrue)[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 的 dispatcherdispatcher 根据设备类型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和backwardbackward里返回对每个输入的梯度。扩展 C/CUDA适合性能敏感、需要底层优化的场景。通过torch.utils.cpp_extension编译自定义 kernel。用Function写自定义算子时最容易出错的是backward的返回值个数和顺序必须和forward的输入一一对应而且只有requires_gradTrue的输入才需要返回梯度其他的返回None。这个规则不遵守反向就会报错或者静默出错。class MyReLU(torch.autograd.Function): staticmethod def forward(ctx, x): ctx.save_for_backward(x) return x.clamp(min0) 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_graphTrue。但要注意保留图会一直占显存用完记得手动释放。我见过有人为了图省事到处加retain_graphTrue结果显存泄漏训练跑一半就崩。理解 PyTorch 内部机制这件事投入产出比其实挺高的。花几个小时把 Tensor、Storage、Autograd、算子这几层的关系理清楚后面遇到问题时定位速度会快很多写自定义算子和做性能优化时也更有底气。我自己的习惯是每学一个新框架都先把它最核心的那两三个抽象搞明白剩下的 API 都是在这上面长出来的。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →