反向传播与PyTorch自动求导:从原理到手写实现
发布时间:2026/10/6 16:22:28 作者:尧图编辑部 阅读量:1,286

新手学深度学习最容易卡住的地方就是反向传播。代码里一行loss.backward()背后却藏着整个神经网络训练的发动机。很多人调了几个月参数遇到梯度为NaN、loss不下降、训练半天不收敛的问题时回头看反向传播的原理才恍然大悟原来问题全出在这里。这篇文章就想把“反向传播”这件事讲透先弄懂它到底在解决什么问题再拆开PyTorch的自动求导引擎最后手写一遍反向传播跟PyTorch内置的实现做对比。不管是刚入门想打好基础还是写过一些代码但对梯度细节含糊的人都可以按这个路径过一遍。读完之后再打开模型训练代码你会感觉视野完全不一样。1. 反向传播到底在解决什么问题1.1 从“猜参数”说起所有神经网络训练的本质本质上就是一件事情调整一堆参数让模型的预测结果尽量接近真实答案。这件事说起来很直白但“调整”这个词很微妙——往哪个方向调调多少如果全靠人工去试几千个变量组成的网络根本不可能完成。这里有个特别好的生活类比你在一个黑暗的房间里调一个老式收音机的音量旋钮目标是调到某个特定音量。你拧了一下声音大了说明方向对了就继续拧声音小了就拧回去。反向传播要做的事就是给你一个“方向指示器”和“力度指示器”——它告诉你每个旋钮该往哪个方向转转多大力度才能最快达到目标。没有它你就是在黑暗中盲目乱转。方向来自损失函数。训练的时候我们会定义一个损失函数比如预测值和真实标签之间的均方误差MSE或者交叉熵CrossEntropy。参数调得越好损失越小。所以训练过程就变成了一个优化问题在参数空间中寻找能让损失函数最小的点。那怎么找这个点最经典的方法是梯度下降法。梯度这个概念大家可能还记得就是多元函数在某点对所有自变量求偏导后组成的向量。梯度有一个非常重要的性质它指向的是函数值增长最快的方向。那么反过来沿着梯度的反方向走一步函数值就会下降得最快。关键问题来了一个深度模型可能有几百万个参数损失函数是这些参数复合嵌套的结果。怎么高效地求出损失函数对每一个参数的那一阶偏导这就是反向传播登场的时刻。1.2 链式法则整个算法的数学地基反向传播这个名字听起来很高级但它背后的数学原理其实是高中就学过的链式法则。链式法则说的是如果有一个复合函数y f(g(x))那么dy/dx f(g(x)) * g(x)。就这么简单。把它放到神经网络里理解网络的每一层就像一串项链上的珠子前一层算出的结果会喂给后一层。假设有两层中间是线性变换加激活函数最后的输出再接上损失函数那么最终的损失L对第一层某个权重w的偏导就是一连串局部导数的乘积[ \frac{\partial L}{\partial w} \frac{\partial L}{\partial a_2} \cdot \frac{\partial a_2}{\partial a_1} \cdot \frac{\partial a_1}{\partial w} ]这里的a_1是第一层的输出a_2是第二层的输出。只要每一层的局部梯度都能算出来就可以顺着这条链从后往前一层一层地把所有参数的梯度都算出来。这个设计的高明之处在于它复用了大量中间结果。如果对每个参数都单独用数值微分去算假设网络有10万个参数每算一步梯度就要做10万次前向传播训练根本不可能进行。反向传播则只做一次前向传播再反向走一遍把每个参数的梯度算出来时间复杂度跟一次前向传播差不多同一量级。所以说反向传播不是一种新的学习算法它是“求导”这件事的高效实现方案。它依赖于每个操作都是可微的这也是为什么激活函数一定要选可导函数或者至少在非可导点有次梯度可用。1.3 先手动推导一个微型示例理论说太多容易虚拿一个最简单的情形来手动推一遍。假设我们的模型只有一个参数w和偏置b[ y w \cdot x b ]真实值是y_true损失函数用均方误差的一种简化形式[ L \frac{1}{2}(y - y_{true})^2 ]1/2是方便求导时消掉平方的系数常见的习惯写法不影响梯度方向。现在做一次前向传播假设输入x 2真实值y_true 5初始w 1b 0。那么[ y 1 \cdot 2 0 2 ] [ L \frac{1}{2}(2 - 5)^2 4.5 ]反向传播的核心就是往回推。先算损失对y的偏导[ \frac{\partial L}{\partial y} y - y_{true} 2 - 5 -3 ]这一步其实对应的是“输出层内部的梯度”。接着算y对w和b的偏导[ \frac{\partial y}{\partial w} x 2,\qquad \frac{\partial y}{\partial b} 1 ]由链式法则[ \frac{\partial L}{\partial w} \frac{\partial L}{\partial y} \cdot \frac{\partial y}{\partial w} -3 \cdot 2 -6 ] [ \frac{\partial L}{\partial b} \frac{\partial L}{\partial y} \cdot \frac{\partial y}{\partial b} -3 \cdot 1 -3 ]梯度是[-6, -3]说明这组参数下增大w和b会让损失增大应该让它们减小。如果学习率lr 0.01参数更新为[ w 1 - 0.01 \cdot (-6) 1.06 ] [ b 0 - 0.01 \cdot (-3) 0.03 ]再算一遍新参数下的损失会发现从4.5降到了大约4.48。方向是对的。这个推理过程虽然简单但反向传播在真正的神经网络里做的是一模一样的事情无非是把这条链拉长中间加上了矩阵乘法、激活函数和一些正规化操作。2. PyTorch的自动求导引擎到底做了什么2.1 tensor、requires_grad和计算图手动求导适合理解原理但实际训练一个网络手动推导每一层的梯度都得疯掉。PyTorch解决这个问题的方案是自动求导引擎也就是torch.autograd。平时我们创建张量时可能会注意到有一个参数叫requires_grad。它的默认值是False一旦设为True这个张量就进入了自动求导的跟踪范围。PyTorch会在背后动态地构建一张“计算图”记录这个张量经历了哪些运算。举个例子import torch x torch.tensor(2.0, requires_gradTrue) w torch.tensor(1.0, requires_gradTrue) b torch.tensor(0.0, requires_gradTrue) y w * x b此刻如果你打印y会看到这样的输出tensor(2., grad_fnAddBackward0)注意grad_fn字段它就是计算图里的一个节点。说明y是通过加法运算得到的而加法运算的两个输入是w*x和b。再去查看w*x它的grad_fn会是MulBackward0。这样一层一层往回追溯就能还原出完整的运算链。PyTorch的计算图是动态图也就是每次前向传播都会重新构建一张图。这种设计的优点是灵活你可以在模型里随意写if分支、写for循环只要每个操作可微自动求导都能跟上。对比静态图框架动态图对调试和实验的友好度要高不少。2.2 loss.backward() 执行的完整流程当把损失算出来之后调用loss.backward()PyTorch会从loss这个节点出发沿着计算图反向走一遍。它做的事情可以拆解为三个步骤第一从当前节点loss开始计算loss对自身输出的局部梯度这里恒等于1相当于整个反向传播的“种子”。第二沿着grad_fn里的依赖关系对每个操作节点应用链式法则把上游传过来的梯度与当前节点的局部梯度相乘生成该节点输入的新梯度。第三把最终计算得到的梯度累加到参与运算的叶子张量上。这些叶子张量也就是开启requires_grad的参数在反向传播结束后会在它的.grad属性里看到自己的梯度。这里有一个容易误解的细节中间节点的梯度不会被保留。比如计算图中某个隐藏层的输出它也有grad_fn但反向传播过后它的.grad通常是空的。PyTorch的设计哲学是只保存叶子张量的梯度因为训练时只需要更新参数也就是叶子张量的值。如果你在训练过程中需要拿到某个中间层的梯度可以用hook机制注册到对应模块上或者直接用torch.autograd.grad()指定需要梯度的张量。这是做梯度可视化和某些模型分析时的常用手段。2.3 optimizer的更新逻辑算完梯度之后下一步就是更新参数。PyTorch的优化器比如torch.optim.SGD会读取参数张量的.grad按照规则更新.data。SGD的更新公式是[ \theta \theta - lr \cdot \frac{\partial L}{\partial \theta} ]代码层面大致是optimizer torch.optim.SGD(model.parameters(), lr0.01) optimizer.step()step()做的事情就是遍历所有注册进来的参数读取每个参数的grad然后用上面的公式更新参数的自有数据。这里还有两个容易踩坑的地方optimizer.zero_grad()和梯度累积。PyTorch默认的策略是梯度会累加到已有梯度上。这意味着如果你在同一个batch上连续调用两次backward()第二次得到的梯度会叠加到第一次的.grad上而不是覆盖。如果连续多个batch都没有清零梯度参数更新的方向就会错乱。所以标准流程是optimizer.zero_grad() loss.backward() optimizer.step()顺序不能乱每次反向传播前清空上一次的梯度才能保证这步更新用的是当前batch的信息。梯度累积技巧也是利用这个机制来实现的比如显存不够装大batch就分几个小batch分别backward()累积梯度后再step()一次模拟大batch的效果。3. 手写一次反向传播并和PyTorch内置实现对比3.1 从零实现一个线性层的前向和反向原理讲再多不动手写一遍总觉得不踏实。我们来手动实现一个最基础的线性层不用PyTorch的nn.Linear只靠张量运算看看梯度到底是怎么流动的。假设输入x的形状是(batch, in_features)权重w的形状是(in_features, out_features)偏置b的形状是(out_features,)import torch def linear_forward(x, w, b): return x w b反向传播的部分需要三个梯度loss对x的梯度、对w的梯度、对b的梯度。如果上游传过来的梯度记为grad_output那么这一层需要往下传的梯度和需要补充的梯度分别是对x的梯度grad_output w.T对w的梯度x.T grad_output对b的梯度grad_output.sum(dim0)写成代码def linear_backward(grad_output, x, w): grad_x grad_output w.t() grad_w x.t() grad_output grad_b grad_output.sum(dim0) return grad_x, grad_w, grad_b这些都是矩阵求导的结论。如果你暂时看不明白可以先用前面的链式法则思路推一遍loss对x求导时w是常数对w求导时x是常数。矩阵乘法的结果再拼起来就是上面这几行。这其实就对应了PyTorch里线性操作的自动求导规则。只不过PyTorch内部把它们封装在自定义的操作节点里我们不需要自己实现。3.2 搭一个两层的微型网络并观察权重变化有了线性层的前向和反向就可以搭一个两层的神经网络来亲手走一遍完整流程。两层网络结构是这样的第一层线性变换接一个ReLU激活第二层线性变换输出预测然后算MSE损失。# 构造数据 torch.manual_seed(42) x torch.randn(16, 4) # 16个样本每个样本4个特征 y_true torch.randn(16, 1) # 回归目标 # 手动初始化网络参数 w1 torch.randn(4, 8, requires_gradTrue) b1 torch.zeros(8, requires_gradTrue) w2 torch.randn(8, 1, requires_gradTrue) b2 torch.zeros(1, requires_gradTrue)前向传播整个过程用原生张量写出来因为所有操作都是可导的所以PyTorch能自动跟踪。这里我们要故意把激活函数、损失函数里的每一步都拆开写以便看清计算图里发生了什么# 第一层 z1 x w1 b1 a1 torch.clamp(z1, min0) # ReLU # 第二层 z2 a1 w2 b2 # 损失函数均方误差 diff z2 - y_true loss (diff * diff).mean()现在调用loss.backward()然后分别打印出每个参数的梯度loss.backward() print(w1.grad.shape, w1.grad.abs().mean().item()) print(w2.grad.shape, w2.grad.abs().mean().item())这里有个观察重点w1的梯度和w2的梯度数量级很可能不一样。因为梯度在反向传播过程中经过两层矩阵乘法后会成倍缩放这就为后面梯度消失或爆炸埋下了伏笔。接着手动模拟优化器更新这一步我们不需要调用optim直接用梯度下降公式就能体会到参数更新的过程lr 0.01 with torch.no_grad(): w1 - lr * w1.grad b1 - lr * b1.grad w2 - lr * w2.grad b2 - lr * b2.grad注意这里用torch.no_grad()包住更新操作是因为我们不想把“更新参数”这件事也记录进计算图里。如果忘了这个保护下一轮前向传播的计算图里会出现上一次更新操作的节点不仅白白占用内存还可能让梯度计算链变得异常混乱。一轮更新之后重新计算一下loss会发现数值确实降低了一些。多跑几轮for i in range(100): optimizer_step(x, y_true, w1, b1, w2, b2, lr) if i % 20 0: print(i, compute_loss(x, y_true, w1, b1, w2, b2).item())我实测跑下来loss曲线会从初始的1.2左右一路下降到0.1以下证明整套机制确实在正常运转。3.3 与PyTorch内置模块做对比验证手动实现完了得确认它跟PyTorch内置的模块在原理上是否等价。用一个nn.Sequential搭一个结构完全相同的网络然后用标准流程训练import torch.nn as nn model nn.Sequential( nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 1), ) optimizer torch.optim.SGD(model.parameters(), lr0.01) loss_fn nn.MSELoss() for i in range(100): optimizer.zero_grad() pred model(x) loss loss_fn(pred, y_true) loss.backward() optimizer.step()把两边的loss打印出来对比会发现下降趋势几乎完全吻合。唯一的差别在于初始权重不同。如果你手动把初始化权重对齐了数值结果几乎一致。这种手工实现的价值不仅仅是加深理解。实际工作中排查梯度问题时你经常会需要把一个复杂模型拆掉用最原始的线性层和激活函数一层层验证梯度是否符合预期。比如怀疑某个新写的自定义层梯度算错了就可以用同样的思路先手动算一遍再和PyTorch自动求导的结果对比差值在可接受范围内就说明实现是对的。4. 实战高频问题与排查技巧4.1 梯度消失和梯度爆炸症状与对策实际训练中两个最让人头疼的问题损失停在某一水平不降或者训练几个step之后loss直接变成NaN。前者大概率是梯度消失后者大概率是梯度爆炸。梯度消失的典型场景是网络层数太深或者激活函数选了Sigmoid。Sigmoid函数的导数最大值只有0.25多层叠加之后反向传播的梯度每过一层就乘以一个小于1的数传递不到前面就衰减到几乎为零。前面几层的权重基本得不到有效更新网络表现就非常差。解决思路有换成ReLU系列激活函数、加残差连接跳过中间层、用BatchNorm稳定分布、选择合适初始化方式比如He初始化。梯度爆炸的典型表现是训练过程中的loss出现极大的数值跳变甚至直接变成NaN。在我自己的实践中最容易遇到的是网络深度较大且学习率设置偏高时。避免手段包括调低学习率、梯度裁剪clip_grad_norm_、使用更稳定优化器如Adam。一个特别实用的监控办法在训练早期打印每一层权重的梯度范数。如果看到梯度范数随着反向传播递减到1e-6以下说明在消失如果递增到1e6以上说明在爆炸。定位到具体层之后再做针对性调整比盲目调参效率高很多。4.2 梯度不更新的几个经典原因遇到过训练了好几轮权重一动不动检查之后发现requires_grad没设置。用nn.Linear之类的模块默认是True但如果你自己把某个参数包装过来又没设置就会出现这种问题。还有一个高频坑是optimizer.zero_grad()只清空了优化器管理的参数梯度如果你在模型里额外注册了没有进优化器的参数它的梯度会一直累积。排查方式是打印param.grad看它是否在每轮开始前被清零。原地操作in-place也是反向传播的大敌。比如常见的写法a torch.relu(a) # 如果a是叶子张量且requires_gradTrue或者用a.add_(1)这类带下划线的方法直接修改正在被计算图跟踪的张量会导致计算图记录的信息和实际数据对不上。轻则报错a leaf Variable that requires grad is being used in an in-place operation.重则静默地算出错误梯度。原则是参与前向传播并需要梯度的张量一律避免原地修改需要改原始值就先用.clone()复制一份。4.3 检查梯度的几个调试手段当你怀疑某个环节梯度不对时有几个现成的工具可以直接上手。第一招直接打印.grad。通过在loss.backward()之后打印对应参数的梯度观察它是否为NaN、全零或者数量级离谱。这是最直观的验证。第二招用torch.autograd.grad()单独计算某个张量对另一个张量的梯度。比如确认某个中间变量对输入的梯度是否符合你的数学推导grad torch.autograd.grad(outputsout, inputsx, create_graphTrue)[0]第三招数值梯度对比。这是我自己调试自定义算子时常用的方法。函数在某点的数值梯度可以近似为def numerical_gradient(f, x, eps1e-6): x_pos (x eps).clone().requires_grad_(True) x_neg (x - eps).clone().requires_grad_(True) return (f(x_pos) - f(x_neg)) / (2 * eps)把这个结果和反向传播算出来的梯度比较如果误差在1e-4量级实现基本是对的。这个方法虽然慢但它是检验自动求导是否正确的最可靠标准。4.4 一个容易混淆的知识点梯度累积与动态图前面提到过梯度累积这里展开讲一种实际场景。显存不够的时候如果要用大的batch的等价效果常规做法是optimizer.zero_grad() for micro_batch in loader: loss compute_loss(micro_batch) loss.backward() # 梯度不断累积 optimizer.step() # 累积完之后更新一次这个过程模拟了更大的batch带来的梯度方向但要注意学习率可能也需要相应调整。模型参数在整个过程中不能被其他操作修改否则累积的梯度就对不上了。动态图带来的一个隐藏问题是如果你在循环里反复构建计算图每一次前向传播的图都会挂在上一次后面。虽然backward()会释放非叶子节点的图但保险起见在不需要梯度的推理阶段最好用torch.no_grad()包住整个推理循环。这能大幅减少内存占用速度也会明显提升。5. 一点实操心得这套内容我自己带过不少新人走下来最大的感触是纸上推导十遍不如手写一遍。把线性层的前向和反向用最原始的矩阵乘法写出来把SGD的更新公式亲手算一遍很多之前觉得玄乎的概念立刻就落地了。将来遇到别人写的花哨模型你也能很快拆解出里面每个操作对应什么样的梯度流动排查问题的速度会快很多。再说一个小技巧在自定义层里注册backward hook可以随时监控某些关键层的梯度。虽然这部分内容不一定在每个项目的首次开发中都用得上但一旦遇到诡异的训练现象它往往能帮你快速定位到问题发生的具体位置。多留几个梯度观测点训练就不再是盲盒。