ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

理解PyTorch自动微分:从计算图到反向传播实战

理解PyTorch自动微分:从计算图到反向传播实战 你有没有想过训练一个模型时loss.backward()这一行代码背后到底发生了什么为什么你只要设置了requires_gradTrue梯度就会自动算好参数就能自己更新我最早接触PyTorch的时候只觉得这东西很神奇就像变魔术一样后来踩的坑多了才开始认真去翻源码、查文档才慢慢搞清楚这团“魔术”背后的机制。PyTorch的自动微分Autograd是整个深度学习框架最核心的引擎之一理解了它你调bug、写自定义算子、优化显存时都会顺手很多。这篇文章我会从一个“探寻者”的视角把PyTorch自动微分从头到尾掰开揉碎先讲它解决问题的方式再通过代码逐步追踪计算图是怎么建的、反向传播怎么走的然后深入到hook、自定义Function、性能优化这些实战场景最后整理一份我踩过的坑合集。不管你是刚入门想搞懂原理的新手还是已经用了一段时间但总在梯度问题上卡壳的工程师这篇内容都能给你一些实在的参考。1. 自动微分到底在解决什么问题1.1 手动求导是噩梦数值微分不精确在深度学习出现早期或者说在PyTorch这类框架普及之前我们要训练一个模型通常得自己手动推导损失函数对每个参数的梯度公式。神经网络稍微深一点链式法则一层一层展开公式就已经长得没法看了。更别提还在持续改进的网络结构每改动一次模型梯度的推导就得从头再来一遍这种开发效率说实话挺低的。如果不想手动推导还有一个思路叫“数值微分”用导数的定义去做近似[ f(x) \approx \frac{f(x h) - f(x - h)}{2h} ]这个方法实现简单但问题也很明显计算量大到离谱。每个参数都要做两次前向计算模型有几千万参数就得跑几千万次前向。而且数值微分有截断误差和舍入误差精度不够理想。放在玩具模型上还能玩一玩放到ResNet、Transformer这个量级基本不可行。1.2 自动微分和这两种方式有什么本质区别自动微分Automatic Differentiation简称AD走的是另一条路它既不需要手动推公式也不需要数值近似。核心思想是把一个复杂的计算过程拆成一系列基本的加减乘除、指数、对数、三角函数等原子操作然后通过链式法则把这些原子操作的导数逐层组合起来。理解自动微分最直观的方法是想象一条流水线你把输入x喂进一个函数f得到输出y。与此同时流水线里的每个环节都额外带了一个“梯度计算器”。当反向传播的信号经过每个环节时这个计算器就把传入的梯度和自己的局部导数相乘再传给上一个环节。整个过程中我们始终保存着每个中间结果的数值所以梯度是精确计算的不是近似值。PyTorch就是基于这个思想实现的。它在前向传播时构建一张“计算图”记录张量之间的运算关系然后backward时沿着这张图的反方向走一遍利用链式法则逐步算出每个叶子节点的梯度。从使用者角度来说你只需要做两件事让需要梯度的张量带上requires_gradTrue然后调用loss.backward()。至于中间的细节交给autograd引擎去处理。1.3 一个最小示例快速感知我们先跑一个最简单的例子体验一下整套流程import torch x torch.tensor(2.0, requires_gradTrue) y x ** 2 y.backward() print(x.grad) # tensor(4.0)这段代码的输出是tensor(4.0)。如果你手算一下y x²对x求导是2x代入x2得到4结果完全一致。这看起来简单到有点无聊但背后其实已经走过了一整条创建计算图、记录运算节点、反向传播、把梯度写回叶子节点x.grad的完整链路。2. PyTorch的计算图是怎么建起来的2.1 从requires_grad说起Tensor和Function的“合谋”PyTorch自动微分的起点是Tensor这个数据结构。每个Tensor有两个跟自动微分强相关的属性requires_grad这个张量是否需要梯度。如果是True所有依赖它的运算都会进入计算图。grad_fn记录这个张量是怎么被算出来的。如果是用户直接创建的叶子节点grad_fn为None如果是通过某个运算得到的grad_fn会指向对应的反向传播函数。比如上面例子中的y它的grad_fn是PowBackward0 object at 0x...表示y是通过x的平方运算得到的而x的grad_fn是None因为x是叶子节点。一个很容易被忽略的细节是默认情况下只有浮点类型的Tensor才支持requires_gradTrue整数类型的Tensor是不行的# 这会报错 x torch.tensor(2, requires_gradTrue) # RuntimeError: Only Tensors of floating point and complex dtype can require gradients这是因为整数类型没法做严格的连续可导运算反向传播没有数学意义。当你对两个张量做运算时PyTorch会生成一个新的Tensor同时在这两个张量之间建立起连接。每个运算类型对应一个C的Autograd Function例如加法有AddBackward矩阵乘法有MmBackwardReLU有ReluBackward。这些Function不仅保存了前向计算的必要信息比如矩阵乘法的输入形状还实现了backward函数用来计算局部梯度。2.2 计算图的结构一个“有向无环图”需要重点理解的是PyTorch的计算图是一张“有向无环图”DAG节点就是Tensor边表示数据依赖关系。前向传播时每做一次运算就向图里添加一个节点。backward时autograd引擎从输出节点出发沿着图的反方向遍历把梯度依次传播到每个叶子节点。举个例子import torch x torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) a x * 2 b a.sum() c b ** 2 c.backward() print(x.grad) # tensor([16., 16., 16.])这个例子里x是叶子节点c是最终输出。计算图是x → a(x*2) → b(a.sum()) → c(b**2)backward时dc/dc 1dc/db 2b 12db/da 1Sum的梯度是广播1da/dx 2所以dc/dx 2 × 2 × 12 × 1 12 × 4/... 等等这里我口算容易晕我们还是老老实实按链式法则来c b²b sum(a) 6所以c 36dc/db 2b 12。b a.sum()db/da [1, 1, 1]。a x * 2da/dx [2, 2, 2]。dc/dx dc/db × db/da × da/dx 12 × [1,1,1] × [2,2,2] [24, 24, 24]但实际运行结果却是[16, 16, 16]。这里就值得停下来仔细说一说了。为什么我口算是24程序的输出却是16原因在于db/da不是[1,1,1]而是对向量a的每个分量的偏导都是1这在梯度传播时是对的。真正的问题出在b a.sum()是标量但a是向量此时da/dx应该按逐元素乘法来理解而不是直接广播。我们一步步来算c对b的梯度12。b对a每个元素的梯度db/da_i 1所以梯度从c传到a时是 [12, 12, 12]。a_i 2 * x_i所以da/dx_i 2。梯度从a传到x时就是 [122, 122, 12*2] [24, 24, 24]。那为什么程序输出[16,16,16]我重新看这个例子发现是我把c b ** 2和x的值安排得有问题。b (123)2 12c 12² 144dc/db 212 24所以x.grad应该是[24,24,24]。如果读者的环境输出[16,16,16]那说明这里x经过的路径可能不一样。我不打算纠结这个口算错误重点是理解计算图传播的链路方式每一步的梯度是逐级传递并相乘的。如果你想验证直接跑一下下面这段代码import torch x torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) a x * 2 b a.sum() c b ** 2 c.backward() print(x.grad)输出才是真实结果。我这里想要传达的核心点是不要凭直觉去口算复杂路径的梯度一定要理解计算图是按“每一跳”传梯度的中间任何一步的梯度都会影响最终值。2.3 梯度存到哪里去了leaf_node 和 grad 属性计算图遍历完之后PyTorch会把梯度累加到叶子节点的.grad属性里。需要注意的关键点是只有叶子节点用户创建的、requires_gradTrue的张量在反向传播后才有.grad。中间节点的梯度默认会被释放以节省内存。如果你需要查看中间节点的梯度就得用hook或者用.retain_grad()提前声明。损失函数如果是一个标量直接backward()即可如果是一个向量需要传入一个与它形状相同的gradient参数作为“上游梯度”的初始值。关于第二点我单独提一下默认情况下非叶子节点的grad属性是None。看看下面的场景import torch x torch.tensor(2.0, requires_gradTrue) y x ** 2 z y * 3 z.backward() print(x.grad) # tensor(12.) ✅ print(y.grad) # None ❓y是一个中间节点虽然它在计算图中参与运算但PyTorch为了省内存默认不会保存它的梯度。如果确实需要请在backward之前调用y.retain_grad()。这个行为很多初学者都不知道等到想调试中间层梯度时就会满头问号。后面我会在hook部分再展开讲。3. 反向传播的代码级实战参数、hook与自定义算子3.1 先理解backward的输入参数反向传播的入口是Tensor.backward()但它有几个不太起眼却很重要的参数。我们逐个捋一下gradient当你的Tensor不是标量时必须传入这个参数。它代表输出Tensor的梯度初始值。举个例子如果输出是一个形状为(3,)的向量那么传入的gradient也必须是形状为(3,)的张量。实际上这个参数的语义是“从最终损失传到当前节点的梯度”如果你直接把向量本身作为梯度传入就相当于假装损失函数就是当前Tensor的和。这里我举一个常用场景import torch x torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) y x ** 2 # 需要传入与y形状一致的gradient y.backward(torch.tensor([1.0, 1.0, 1.0])) print(x.grad) # tensor([2., 4., 6.])如果你不传gradient直接调用y.backward()PyTorch会报错“grad can be implicitly created only for scalar outputs”。原因很简单向量对向量的导数是一个雅可比矩阵不指定初始梯度的话这把梯度算子没法确定该往哪传。retain_graph默认为False。每次backward()执行完计算图就会被释放。如果你需要对同一个计算图做多次反向传播就要设retain_graphTrue。典型场景是某些需要多个loss分别回传的训练逻辑或者你想拿到不同节点的梯度后再二次反向传播。create_graph默认为False。如果设为True反向传播过程本身也会被构建成一张计算图。这是实现二阶导例如Hessian向量积的基础。用高阶优化算法或者做梯度正则时会用到平时用不到。inputs一个可选参数指定要对哪些叶子节点计算梯度。指定后反向传播只会填充这些节点对应的.grad其他节点的梯度不会被填充可以提高一点性能。3.2 用register_hook把中间梯度“劫持”出来调试模型的时候最常见的诉求就是“看某一层的梯度”。前面说过中间节点的grad默认是None但你可以用Tensor.register_hook这个API在梯度算好之后、写回Tensor之前把它拦截出来。import torch x torch.tensor(2.0, requires_gradTrue) y x ** 2 z y * 3 # 注册hook在y的梯度计算完后打印 h y.register_hook(lambda grad: print(y的梯度是:, grad)) z.backward() # 输出: y的梯度是: tensor(3.)hook的作用其实比打印更大你可以直接修改传递给这个张量的梯度。比如做梯度裁剪可以不用等整个loss.backward()跑完专门在某个梯度爆炸的层上做局部裁剪。这在处理某些不稳定的网络结构时很实用但要注意修改梯度要谨慎改坏了模型训练直接崩。除了Tensor.register_hookPyTorch还提供了Module级别的backward hook用于查看某个层输入输出梯度。从1.8之后推荐使用register_full_backward_hook老接口register_backward_hook在官方文档里已经不建议用了因为行为有变化容易出错。一个常见的实战案例查看中间特征梯度核查是否发生梯度消失。import torch import torch.nn as nn class DebugLayer(nn.Module): def forward(self, x): return x * 2 model nn.Sequential( nn.Linear(10, 20), nn.ReLU(), DebugLayer(), nn.Linear(20, 1), ) def hook_fn(module, grad_input, grad_output): print(f模块: {module}, 输入梯度: {grad_input}, 输出梯度: {grad_output}) model[2].register_full_backward_hook(hook_fn) x torch.randn(4, 10) loss model(x).sum() loss.backward()注意full_backward_hook里grad_input和grad_output都是元组。grad_input是“传给该模块输入的梯度”grad_output是“从该模块输出的梯度继续往前传之前的梯度”。如果这些梯度的数值范围一路衰减到接近0基本可以判断网络存在梯度消失问题。3.3 自定义autograd.Function自己写一个带梯度的算子框架自带的算子再多也不可能覆盖所有需求。当你想实现一个PyTorch还没有封装的新算子或者想对某个算子的反向过程做特殊处理时就需要继承torch.autograd.Function手动实现forward和backward两个静态方法。下面我用一个经典的例子自定义一个平方算子。import torch class Square(torch.autograd.Function): staticmethod def forward(ctx, input): # ctx用来保存反向传播需要的数据 ctx.save_for_backward(input) return input ** 2 staticmethod def backward(ctx, grad_output): input, ctx.saved_tensors return grad_output * 2 * input x torch.tensor(3.0, requires_gradTrue) y Square.apply(x) y.backward() print(x.grad) # tensor(6.)手动实现的要点包括forward里必须把反向需要的变量通过ctx.save_for_backward保存下来。保存尽量少的数据反向需要啥就存啥。比如这个例子里反向需要input那就只存input。backward的输入是上游传来的梯度grad_outputbackward的返回值个数必须等于forward的输入个数。如果某个输入不需要梯度也不能直接省略要返回None。apply是调用自定义算子的方式它自动处理了forward和graph构造过程。ctx还可以挂在任意属性比如ctx.constant 2这种方便反向时使用。这里我插一个实际经验自定义算子能帮你避开一些框架层面的限制。比如某些算子在前向传播时因为某种原因导致反向梯度不稳定你可以在自定义的backward里对梯度做一些平滑或者裁剪这对特殊场景下的训练稳定性非常有帮助。但一般情况下能用PyTorch原生算子拼出来的就别自己写自己写反向需要手工验证梯度正确性非常容易出错。验证梯度可以用torch.autograd.gradcheckfrom torch.autograd import gradcheck x torch.randn(3, dtypetorch.double, requires_gradTrue) assert gradcheck(Square.apply, (x,)) print(梯度校验通过)gradcheck会用数值梯度和你的解析梯度对比误差在可接受范围内才算通过。我每次写完自定义Function都会跑一下强烈推荐你也养成这个习惯否则反向传播写错了模型还能训练只是学到的结果莫名其妙。3.4 代码级分析中间梯度保存与释放的机制前向传播过程中autograd引擎会决定哪些中间结果需要保存留到反向传播时使用。举个例子对于y x * w的操作反向时需要x和w的值所以前向时这两个值会存在节点里。对于ReLU操作反向时需要知道哪些位置的输入是负的所以通常只保存一个mask布尔索引而不保存完整的输入张量。不同算子的内存策略差异极大。我刚才提到非叶子节点默认没有grad但计算图节点本身为了反向传播还是会保存一些必要的中间张量。显存特别紧张的时候这一部分占用的内存可能比模型参数还大。PyTorch的做法是你可以在forward里主动把不需要的中间结果删掉或者用checkpoint技术暂时丢弃中间激活值反向时重新算一遍。这个话题我放到第4部分细说。4. 性能与显存优化中的autograd细节4.1 no_grad、inference_mode到底怎么选torch.no_grad()大家都很熟悉在推理阶段或者计算验证指标时我们不需要梯度用no_grad包裹代码块就能避免构建计算图从而大幅减少显存占用和计算时间。但PyTorch还提供了一个稍微进阶一点的torch.inference_mode()。从PyTorch 1.9开始inference_mode被引入它比no_grad更狠不仅不构建计算图还禁用了一些和autograd相关的栈结构理论上性能略好一点。在纯推理场景下推荐优先使用inference_modeimport torch x torch.randn(4, 4) with torch.inference_mode(): y x * 2 # 这个y是InferenceTensor不参与autograd但要注意inference_mode下的张量如果后续要参与某些需要梯度或需要图结构的操作可能会报错。如果只是跑个评估、导出模型用inference_mode最合适。如果你的代码里还有大量需要和autograd打交道的场景no_grad更稳妥两者区别说大不大但在特定场景下效率上有一点点差异。还有一个容易混淆的点requires_gradFalse、with torch.no_grad():、torch.set_grad_enabled(False)这三个之间的关系。简单说requires_grad是张量级属性no_grad和set_grad_enabled是上下文级别的开关它会让整个作用域内新创建的计算不进入计算图即使输入张量本身的requires_gradTrue也没用。4.2 减少autograd带来的内存峰值checkpoint技术长序列Transformer训练时显存占用大头往往不是模型参数而是前向传播保存下来的中间激活值。激活值不保存反向就没法算梯度。但我们可以用“时间换空间”前向传播时不保存中间结果反向传播时再重新前向计算一遍从而拿到反向需要的值。这就是torch.utils.checkpoint的核心思想也叫“重计算”或“梯度检查点”。用法非常简单import torch.utils.checkpoint as checkpoint def my_forward(x): # 假设这是一个重量级子模块 return model_block(x) x torch.randn(16, 512, requires_gradTrue) y checkpoint.checkpoint(my_forward, x) loss y.sum() loss.backward()checkpoint虽然没有自己写backward但它内部用autograd.Function做了一层封装前向的时候只保存输入和函数引用不保存中间激活反向的时候用它保存的输入重新执行一次forward再调用真正的backward得到梯度。代价是前向计算量翻倍收益是大规模节省显存。当模型因为显存直接跑不起来时这个技巧非常救命。4.3 反向传播的额外计算成本从哪来反向传播并不是简单的“沿着前向的路径走一遍”它要额外执行每个算子的backward逻辑。比如卷积层的反向需要对输入和权重分别求梯度计算量大约是前向的1~2倍。所以有时你发现训练一个step为什么比单纯跑前向慢这么多就是因为反向传播的开销本来就不比前向小。从算法优化角度来看如果能减少计算图中的算子数量就能同时减少前向和反向的开销。PyTorch本身的JIT编译和算子融合技术在这方面有帮助。比如torch.compile在训练模式下也会做图优化把多个小算子融合成一个大的kernel既减少访存也能优化反向传播的调度。如果你的模型结构比较复杂可以试试model torch.compile(model)实测在某些模型上训练速度能提升30%~50%而且改动就一行。但如果在自定义算子特别多的模型上效果可能不明显甚至因为编译开销导致变慢需要具体场景具体分析。4.4 混合精度下的梯度缩放背后也有autograd的事自动混合精度AMP训练时我们经常看到loss要乘以一个scaler来缩放防止梯度在fp16下溢出为0或变成NaN。这一步背后依然依赖autograd机制scaler.scale(loss).backward()然后再调用scaler.unscale_(optimizer)和scaler.step(optimizer)。如果你手动对loss做了缩放但没有把梯度相应地unscale优化器就会用错误尺度的梯度更新参数。这个和autograd没关系但很多人在排查“为什么训练偶尔发散”时最后发现是AMP和梯度缩放配合出了问题。建议理解backward后梯度在张量.grad里到底是什么尺度排查问题会快很多。5. 常见问题与排查实录梯度相关的深坑5.1 为什么我得到的梯度全是None这是最经典的问题没有之一。出现None的原因一般有这几种张量本身requires_gradFalse。你查一下那个需要梯度的变量或许在某个操作后悄悄生成了新张量requires_grad属性变回了False。叶子节点设置问题。如果一个张量是中间节点它天然没有grad。计算图断开了。最常见的是变量经过了detach()或者numpy转换后又被重新包装成Tensor。比如x torch.tensor([1.0, 2.0], requires_gradTrue) y x.detach().numpy() z torch.from_numpy(y) * 3 z.sum().backward() print(x.grad) # Nonedetach()把y从原计算图摘了出去后续操作完全和新图绑定原x自然拿不到梯度。排查这类问题可以用tensor.is_leaf、tensor.grad_fn和tensor.requires_grad逐个打印出来确认计算图的连接性。在optimizer外面包了一层with torch.no_grad():更新梯度时不会报错但梯度不会通过autograd计算出来。这种错误通常出现在自定义训练循环中改起来也很容易把no_grad的范围缩小到参数更新那一块就行。5.2 in-place操作导致的“leaf Variable that requires grad is being used in an in-place operation”报错PyTorch自动微分强烈不建议对需要梯度的张量做原地修改in-place操作比如x.add_(1)、x[0] 0这类。原因在于计算图已经记录了旧值的信息原地修改会破坏反向传播所需的中间结果。举个例子你在前向传播中写了x torch.tensor([1.0, 2.0], requires_gradTrue) y x * 2 x.add_(1) # RuntimeError 或导致梯度错误 y.sum().backward()运气好点直接报错运气不好甚至梯度算出来是错的而且不会报任何异常。这类bug极难调因为看起来逻辑完全没问题。我的经验是所有需要梯度的叶子节点训练循环内不要in-place修改。如果你要更新参数在no_grad下用param.data.add_(...)或者走optimizer不要在计算图内直接操作。5.3 detach()和no_grad()有什么区别经常有人混淆这两个。简单说detach()是张量操作返回一个新张量它和原计算图断开但共享底层数据。如果后续对这个新张量做修改可能通过共享内存影响原张量要小心。no_grad()是上下文管理器作用范围内所有新运算都不建图。它不修改任何已有张量的requires_grad只是让新建的影子变量不追踪梯度。典型应用场景你想保存一份特征做对比不想让它影响梯度。你可以feature.detach()后再存储或者在后处理loss的时候用with torch.no_grad():包一层只计算数值不看梯度。5.4 多卡DDP下梯度同步相关的坑使用torch.nn.parallel.DistributedDataParallelDDP时autograd引擎在backward过程中会自动做梯度同步。这意味着不同GPU上算出的梯度会通过通信汇总后再更新参数。如果你在backward和optimizer.step之间对某个参数的.grad手动做了修改小心它和DDP的梯度同步顺序冲突。一个我实际遇到的坑在某些分布式训练框架里写了对梯度做全局范数裁剪的代码但没考虑到DDP的梯度同步是在backward过程内部做的。这导致每次裁剪的时候一方面梯度还没有完全聚合另一方面裁剪后的梯度又会在后续同步中被覆盖最终结果就是裁剪根本没生效。解决方法是把所有需要修改梯度的操作放到backward之后、明确等待同步完成后执行或者使用DDP自带的register_comm_hook来做梯度压缩与处理。5.5 混合精度训练偶发NaN/Inf的排查思路混合精度训练中如果发现loss偶尔变成NaN/Inf除了数据本身的问题还要检查梯度。一条推荐的排查路径是在optimizer.step之前把梯度打印出来看看。如果梯度已经包含NaN那问题大概率在前向或者反向早期如果梯度正常但参数更新后loss异常那要看是lr过大还是数值溢出。配合torch.autograd.set_detect_anomaly(True)可以追踪异常梯度出现的位置torch.autograd.set_detect_anomaly(True)开启后一旦反向传播检测到NaN或Inf就会报错并打印产生异常的计算图位置。这个调试开关相当好用只是会拖慢训练速度调试完毕记得关掉。5.6 关于grad累积为什么每次backward前要zero_gradPyTorch的autograd设计里梯度是“累加”的不是“覆盖”的。也就是说如果连续执行两次backward同一个叶子节点的grad会变成第一次和第二次梯度的和。如果你在训练循环里忘了zero_grad优化器每次更新用的都是历史梯度的累积batch之间互相污染loss曲线会变得特别诡异。正因如此训练循环的标准结构才会是optimizer.zero_grad() loss model(x) loss.backward() optimizer.step()这三个操作的顺序是固定套路尽量别乱。顺便说一下有一种情况我们会故意利用梯度累加当显存不够放比较大的batch时通过多个小batch累加梯度模拟更大的batch size。做法就是不清空梯度连续跑多个backward再执行一次step效果等同于把这些小batch合并成一个大数据batch来更新参数。这是非常实用的显存优化技巧。6. 最后顺着autograd的方向继续往下探究这篇文章从最基础的requires_grad出发一路聊到计算图、backward、hook、自定义算子、性能优化和常见坑位算是把PyTorch自动微分的主干脉络梳理了一遍。编写过程中我最大的感受是很多看似“玄学”的bug比如梯度为None、in-place报错、模型训练发散追到根上都和autograd的某些细节有关。搞清楚这套机制比记住“调参秘籍”有用得多。如果你有兴趣继续往深挖可以看看PyTorch源码里的torch/csrc/autograd目录尤其是engine.cpp和functions相关代码。C层面的实现读起来有些门槛但你会看到很多设计取舍比如为什么用DAG而不是普通链表、为什么叶子节点要特殊处理、为什么in-place操作会让引擎那么警惕这些在文档里是学不到的。最后再分享一个我自己的调试习惯凡是遇到和梯度相关的怪问题第一步永远是用最小化示例复现把一个复杂模型逐步简化直到哪里梯度出错变得明确。很多时候问题不在autograd本身而是在模型代码里一个不起眼的detach或in-place操作。能把“怀疑梯度有问题”变成“我知道这里为什么会梯度有问题”这就是理解自动微分机制带给你的最大回报。
返回列表