新手学深度学习,最容易卡住的地方就是反向传播。代码里一行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 = 1,b = 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_grad=True) w = torch.tensor(1.0, requires_grad=True) b = torch.tensor(0.0, requires_grad=True) y = w * x + b此刻如果你打印y,会看到这样的输出:
tensor(2., grad_fn=<AddBackward0>)注意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(), lr=0.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(dim=0)
写成代码:
def linear_backward(grad_output, x, w): grad_x = grad_output @ w.t() grad_w = x.t() @ grad_output grad_b = grad_output.sum(dim=0) 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_grad=True) b1 = torch.zeros(8, requires_grad=True) w2 = torch.randn(8, 1, requires_grad=True) b2 = torch.zeros(1, requires_grad=True)前向传播整个过程用原生张量写出来,因为所有操作都是可导的,所以PyTorch能自动跟踪。这里我们要故意把激活函数、损失函数里的每一步都拆开写,以便看清计算图里发生了什么:
# 第一层 z1 = x @ w1 + b1 a1 = torch.clamp(z1, min=0) # 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(), lr=0.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_grad=True或者用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(outputs=out, inputs=x, create_graph=True)[0]第三招,数值梯度对比。这是我自己调试自定义算子时常用的方法。函数在某点的数值梯度可以近似为:
def numerical_gradient(f, x, eps=1e-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,可以随时监控某些关键层的梯度。虽然这部分内容不一定在每个项目的首次开发中都用得上,但一旦遇到诡异的训练现象,它往往能帮你快速定位到问题发生的具体位置。多留几个梯度观测点,训练就不再是盲盒。