news 2026/10/6 16:22:11

反向传播与PyTorch自动求导:从原理到手写实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
反向传播与PyTorch自动求导:从原理到手写实现

新手学深度学习,最容易卡住的地方就是反向传播。代码里一行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,可以随时监控某些关键层的梯度。虽然这部分内容不一定在每个项目的首次开发中都用得上,但一旦遇到诡异的训练现象,它往往能帮你快速定位到问题发生的具体位置。多留几个梯度观测点,训练就不再是盲盒。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/6 16:22:10

C#与SQL Server网上书店系统:三层架构与并发扣库存实战

简介&#xff1a;基于C#与Sql Server的网上书店管理系统&#xff0c;是一份适用于ASP.NET课程设计或毕业设计场景的完整项目资料&#xff0c;面向需要学习B/S结构开发、数据库设计及前后台交互的读者。系统采用B/S架构&#xff0c;前台提供用户注册登录、商品浏览、购物车、订单…

作者头像 李华
网站建设 2026/10/6 16:22:10

C++数据结构学习与期末复习:北理工资源全解析

简介&#xff1a;北理工2020年《数据结构》课程资料包&#xff0c;面向正在学习C与数据结构的学生&#xff0c;覆盖从基础概念到算法实现的全流程。压缩包共65个文件&#xff0c;包含29个cpp源代码、9个ppt课件、16个doc和5个docx文档&#xff0c;另有5个pdf试卷及1个pptx讲义&…

作者头像 李华
网站建设 2026/10/6 16:22:06

OFDM循环平稳检测与协作频谱感知:从原理到工程避坑

简介&#xff1a;面向OFDM通信与认知无线电频谱感知研究的MATLAB仿真源码包&#xff0c;适用于高校无线通信课程设计、论文仿真及入门学习者。针对阴影和深度衰落下单节点感知结果不可靠的问题&#xff0c;代码覆盖循环平稳特征检测、能量检测以及协作频谱感知&#xff0c;可在…

作者头像 李华
网站建设 2026/10/6 16:20:58

15个Git核心命令,覆盖日常开发全流程

前一阵带新同事熟悉项目&#xff0c;他一边翻Git命令手册一边叹气&#xff0c;说命令太多&#xff0c;背了又忘&#xff0c;干脆继续用图形界面。我给的建议是&#xff1a;别背&#xff0c;就把15个命令用熟&#xff0c;足够应付日常开发了。这篇聊聊我筛选出的这15个Git核心命…

作者头像 李华
网站建设 2026/10/6 16:20:04

JSP百货供应链管理系统课设:从环境搭建到答辩的完整指南

简介&#xff1a;这份资源是面向计算机专业学生与Java Web初学者的一套百货中心供应链管理系统完整毕业设计资料&#xff0c;包含可运行的JSP项目源码、数据库脚本与WORD论文文档&#xff0c;适合用作课程设计、毕业设计参考或供应链管理系统的学习案例。压缩包共10个文件&…

作者头像 李华
网站建设 2026/10/6 16:19:19

方维3.4 P2P借贷系统源码解析:PHP交易系统的部署与改造

简介&#xff1a;方维3.4专业P2P网络贷款借贷系统是一套可直接部署的PHP源码包&#xff0c;面向需要搭建网络借贷、投资理财平台的站长、开发者及中小团队&#xff0c;既可商用二次开发&#xff0c;也适合学习经典P2P系统的前后端结构。完整包内共2000个文件&#xff0c;其中89…

作者头像 李华