Skip to content

s06 反向传播与链式法则

WARNING

🧪 Beta公测版本提示:教程主体已完成,正在优化细节,欢迎大家提Issue反馈问题或建议。

每个节点只关心自己的局部导数 —— 揭开 autograd 的魔法


一、核心问题:一个参数的微小变化如何影响损失?

上一节 s05 计算图与前向传播 把输入推到了预测 y^。训练要做的事更具体:拧动权重,让 y^ 贴近标签 y。反向传播回答的就是——当前这一步,每个权重该对「错」负多少责。

拆链式法则之前,必须先讲清两件事:损失 L 是什么、MSE 怎么来的,以及为什么求的是损失对权重的梯度(而不是对输入、也不是停在「对预测的梯度」)。

1.1 损失:把「错得有多离谱」压成一个数

网络里有成千上万个权重,不可能靠肉眼逐个改。必须先有一个标量分数 L

  • L 小:预测整体接近标签
  • L 大:偏差大,模型「错得离谱」

训练的定义就是找参数让这个分数最小:

θ=argminθL(θ)

s05 里 L=(y^,y) 还只是占位符。下面把本章(以及 demo.py)用的那一种——均方误差 MSE——从残差一步步推出来。

1.2 MSE 是什么、怎么来的

从最朴素的问题开始:预测和标签差多少?

e=y^y

这个差叫做残差(residual)。为什么不能直接把残差(或残差的平均)当成损失?

  1. 正负会抵消。 一个样本偏高 3、另一个偏低 3,平均残差是 0,模型看起来「完美」,其实两个都错了。
  2. 我们关心的是「差了多少」,不是「偏高还是偏低的代数和」。

自然的补丁是取绝对值 |y^y|(MAE)。它能用,但在 e=0不可导,梯度下降会在「刚好命中」附近卡住;而且大错误和小错误被一视同仁。

再进一步:把残差平方。单样本损失(本章公式,带 12)是

(y^,y)=12(y^y)2

N 个样本再取平均,就得到均方误差(Mean Squared Error, MSE)

L=1Ni=1N12(y^iyi)2

「均」= 对样本平均,「方」= 残差平方。平方带来四件我们真正需要的性质:

  1. 永远 0,且仅当 y^=y 时为 0——完美命中时损失触底。
  2. 处处光滑可导。反向传播的第一枪就是 /y^,没有导数就传不回去。
  3. 大偏差惩罚更重。 误差 10 的平方是 100,误差 1 的平方是 1。模型会被逼着先修离谱的错。
  4. 统计来源,不是拍脑袋。 若观测 = 真值 + 噪声,且噪声 εN(0,σ2),那么最大似然估计(MLE)恰好等价于最小化平方误差。高斯噪声下,「最好的拟合」就是 MSE。

MSE 与 MAE 的几何对比,线性回归 第 3 节写得更细。本章要抓住的是:MSE 给了反向传播一个可导的标量起点

12 从哪来? 求导时 212 约掉:

y^=y^y

没有 12 时梯度是 2(y^y),只差一个常数,极小值的位置不变,学习率可以吸收这个倍数。本章公式带 12,是为了让链式法则第一项干净。demo.py 里写的是 (y^y)2,没除 N、也没乘 12——优化方向相同。

分类任务常用交叉熵,公式不同,角色相同:仍然是一个标量 L,仍然要求 L/w

对 MSE,反向传播的第一枪立刻能写出来:

Ly^=y^y

含义直白:预测比标签高多少,就把这么大的「不满」往回传。y^ 偏高,这项为正,后面会推动权重把预测压下来。

1.3 为什么求的是「损失对权重」的梯度?

有了 L,对谁求导?计算图里出现过三类量:

变量训练时能不能改对它求梯度有没有用
输入 x一般不改,那是数据L/x 另有用途(显著性、对抗样本),不是训练步骤
标签 y不当成旋钮
权重 w、偏置 b唯一能拧的旋钮训练要的就是 L/wL/b

所以目标不是停在「损失对预测的梯度」,而是把这笔账算到每个权重头上。

为什么是梯度(一阶导数),而不是别的?把 L 看成 w 的地形:

  • L/w>0:在这一点增大 wL 会升高 → 应该减小 w
  • L/w<0:增大 w 会让 L 降低 → 应该增大 w

梯度指向 L 上升最快的方向;我们要下山,所以反着走:

wwαLw

α 是学习率,控制这一步迈多大。这就是梯度下降。s05 写过这行更新式;反向传播的职责是把式子里的 L/w 对每一个参数都算出来。

只知道 L/y^ 为什么不够? 它只说「输出偏高了 0.3」,没说是 w3 还是 w17 造成的,每个该改多少。必须再乘「这个权重怎样影响输出」:

Lw=Ly^y^w

这就是下一节的链式法则。L/w 的含义:当前这个权重,对总损失要负多大的责

用手算钉死一次(对应下图右侧)。模型 y^=wx,取 x=2y=1,于是 L=12(2w1)2。若当前 w=1,则 y^=2,偏高 1

Ly^=y^y=1,y^w=x=2,Lw=12=2>0

梯度为正,减小 w。最优解 w=1/2,此时 y^=y=1。符号和大小都对得上:预测偏高,而且 x=2 把这份不满放大了一倍——这个 w 对输出的影响力就是 2

MSE 把残差变成可导的标量损失;L(w) 上梯度指上坡,更新走下坡

三件事的分工:损失函数规定「什么叫错」(MSE 或交叉熵);反向传播高效算出 θL;优化器决定「知道责任之后怎么改」(SGD、Adam)。本章做中间那一步,但必须从 MSE 这个起点出发。

现在可以问反向传播的核心问题了。L 并不直接写成 w 的式子——中间隔着 za。若让 w 增加一个微小的 Δw,损失会变化多少?

Lw=?

链式法则就是把这条间接依赖拆开的数学工具。

单神经元链式法则详解


二、链式法则:从简单神经元开始

考虑一个最小化的神经元模型:

z=wx+ba=ϕ(z)L=(a,y)

我们想知道 L/w。根据链式法则,将 Lw 的依赖沿着中间变量 az 拆开:

Lw=Laazzw.

反向传播:链式法则从输出往回乘

图解说明:前向算出 z,a,L;反向从 L/a 起,每层只乘自己的局部雅可比,一直乘到 L/W

逐步推导:三层小神经元上的 L/w(点击展开)

z=wx+ba=ϕ(z)L=(a,y)。链式:

Lw=Laazzw.

最后一项 z/w=x。中间 a/z=ϕ(z)。最外层 L/a 由你选的损失决定(MSE 时是 ay)。三者相乘就是一次反向传播。多层只是把这段重复:每一层缓存前向的 z,反向时用后一层送来的 L/a 乘本地 ϕ 再乘 x(或 W)。这就是「反向」:信息从损失往输入流,参数就地更新。

逐项拆解:

  • 第一项 La:损失对激活输出的梯度。第 1.2 节已经推出:MSE 下 (a,y)=12(ay)2,所以这一项就是 ay(「预测偏了多少」)。交叉熵的公式不同,但同样是一个标量 L 对输出的导数,后面的链式法则不变。
  • 第二项 az=ϕ(z):激活函数的导数。这一项在 z 处取值,所以前向传播时必须存储 z
  • 第三项 zw=x:因为 z=wx+b,对 w 求偏导就是 x。类似地,z/b=1

把它们乘起来:

Lw=La损失梯度ϕ(z)激活导数x输入

对偏置的梯度更简单:

Lb=Laϕ(z)1=Laϕ(z)

直观理解w 的梯度 = "损失对你的输出有多不满" × "激活函数在当前点的斜率" × "这个参数对应的输入值的大小"。如果输入 x 很大,那么 w 的梯度也大,意味着这个 w 对最终结果影响大,需要更大的调整。

demo.py 用手算表达式 f=(x+y)z 对账框架的 backward。再对照单神经元:y^=wxx=2y=1w=1L/w=2,见第 1 节。卡点:fan-out 必须 grad +=;忘了 zero_grad 会把上一步梯度叠进来。


三、为什么叫"反向"传播?

"反向"是相对于前向传播而言的:

  • 前向传播:数据从输入流向输出(xzaL),这是计算预测值的过程。
  • 反向传播:梯度从损失流回输入(Lazw),这是计算梯度信息的过程。

反向传播的计算顺序恰好与前向相反。这不是人为规定,而是数学必然——链式法则的每一项都依赖"后面"的计算结果:

  • 要算 L/z,需要先知道 L/aϕ(z)
  • 要算 L/w,需要先知道 L/zx

所以算法的自然流程是:先一路前向算到损失,再一路反向把梯度传回去。这也是为什么前向传播时要把中间值都存起来——反向时会逐个用到。


四、计算图视角:每个节点只关心自己

反向传播真正的"魔法"在于:网络再深,每个节点只需要实现自己的局部导数规则

考虑计算图中的一个乘法节点:

u=pq

当反向传播时,这个节点收到"上游"传来的梯度 L/u。它需要做的是:算出梯度如何分配给自己的两个输入 pq

因为 u/p=qu/q=p,所以:

Lp=Luup=LuqLq=Luuq=Lup

注意一个有趣的现象:乘法门的梯度分配是"交换"的——p 的梯度乘的是 q 的值,q 的梯度乘的是 p 的值。这在英文文献中常被称为"gradient switcheroo"。

再看加法节点:

u=p+q

因为 u/p=1u/q=1

Lp=Lu,Lq=Lu

加法门像是一个"梯度分发器"——上游梯度原样复制给每个输入。这就是为什么在 x += y 操作中,梯度是累积的。

常见计算图节点的反向规则


五、常见操作的局部梯度规则

以下是计算图中常见操作的反向传播规则,可以当作参考卡片使用:

前向操作局部导数反向规则
u=p+qu/p=1, u/q=1梯度原样传递:L/p=L/u
u=pqu/p=1, u/q=1q 的梯度取反
u=pqu/p=q, u/q=p梯度交换(乘以对方的
u=p/qu/p=1/q, u/q=p/q2注意 q 的梯度含负号
u=pku/p=kpk1相当于 kpk1 的乘法
u=max(0,p) (ReLU)u/p=1 if p>0 else 0像一个"门"——根据前向时的 p 值决定梯度是否通过
u=σ(p) (sigmoid)u/p=u(1u)前向输出即可算出导数,无需知道 p
u=tanh(p)u/p=1u2同样可以用输出直接算导数
u=exp(p)u/p=u梯度等于自身的值

关键思想:每种操作都是独立的"乐高积木"。深度学习框架(PyTorch、TensorFlow、JAX)在底层为每种操作都实现了 forward()backward() 两个函数,然后把它们按计算图拼起来。当你调用 .backward() 时,框架不过是在按拓扑序的逆序逐个调用这些节点的 backward()


六、梯度累积:多路径的 Fan-Out

一个变量可能在计算图中被多次使用(fan-out)。例如:

x ──┬──→ [×2] ──→ u ──┐
    │                   ├──→ [u·v] ──→ L
    └──→ [+3] ──→ v ──┘

x 同时影响了 u(通过 ×2)和 v(通过 +3),而 uv 又共同影响了 L。此时,x 的梯度来自两条路径的梯度之和

Lx=Luux+Lvvx

这是多元微积分中全导数的链式法则:

Lh=iLuiuih

其中 ui 是计算图中以 h 为输入的所有节点。

在代码实现中,这就是为什么梯度要累加grad += ...),而不是直接赋值(grad = ...)。PyTorch 中的 .backward() 默认就是累加模式,所以在每次反向传播前需要调用 optimizer.zero_grad() 来清零。

Fan-Out:多路径梯度的求和


七、完整示例:用手算理解反向传播

让我们通过一个具体例子,手动做一遍前向和反向计算。考虑表达式:

f(x,y,z)=(x+y)×z

给定具体值:x=2,y=3,z=4

前向传播(Forward Pass)

  1. u1=x+y=2+3=5
  2. u2=u1×z=5×4=20

所以 f(2,3,4)=20

反向传播(Backward Pass)

我们要计算 f/xf/yf/z

从输出端开始(设 L=f,即 L/u2=1):

步骤 1:穿过乘法门 u2=u1×z

Lu1=Lu2z=1×4=4Lz=Lu2u1=1×5=5

步骤 2:穿过加法门 u1=x+y

Lx=Lu11=4Ly=Lu11=4

验证

f/x=z=4(因为 f=(x+y)zf/x=z)。与我们算的 4 一致。

f/y=z=4。与我们算的 4 一致。

f/z=x+y=5。与我们算的 5 一致。

这就是反向传播的全貌——一个具体的、可复现的过程。所有深度学习框架的 backward() 不过是在更大的计算图上做同样的事情。

反向传播分步数值示例


八、反向传播的复杂度分析

反向传播的一个重要特性是它与前向传播有相同的计算复杂度(量级上)。具体来说:

  • 前向传播:遍历一次计算图,每个节点计算一次
  • 反向传播:逆向遍历一次计算图,每个节点计算一次

两者都是 O(N) 的时间复杂度,其中 N 是计算图中节点的数量。这意味着训练时间大约是推理时间的两倍——一次前向,一次反向。

空间复杂度则更高,因为反向传播需要存储前向传播的所有中间值。对于有 N 个节点的计算图,空间复杂度也是 O(N)。这在大模型训练中是显存的主要消耗来源之一。


九、从手动求导到自动微分

人类手动求导的方式是:写出整个函数的符号表达式,然后对每个参数求导。对于浅层网络这还可操作,但面对上千层的网络,写出一个包含百万参数的巨大导数表达式是完全不可能的。

自动微分(Automatic Differentiation, AD)的解决方式是:

  1. 把复杂函数拆成基本操作(加、减、乘、除、指数、log、sin 等)。
  2. 为每个基本操作手工实现一个 "局部 backward"——就像我们上面整理的规则卡片。
  3. 前向计算时,记录操作顺序和中间值。
  4. 反向时,按逆序依次调用每个操作的 backward。

这种方法叫做反向模式自动微分(Reverse-mode AD),是深度学习框架的核心引擎。它的关键优势是:不管函数多复杂,只要它能被分解为基本操作,就能自动求出所有参数的梯度——而且计算量和前向是同量级的。

下一节 s07 多层网络的矩阵反传 将在矩阵层面推导完整的反向传播公式(包括 δ 递推关系),并实现一个完整的训练循环。


十、本节小结

概念一句话
MSE12(y^y)2:把残差平方成可导的非负标量;12 只为让导数等于 y^y
L/w只有权重能改;梯度指向上坡,更新 wwαL/w 走下坡
链式法则将间接依赖的梯度拆成局部导数连乘
反向顺序从损失往回算,自然地对应链式法则的计算顺序
局部梯度每个节点只需知道自己的导数规则,不关心网络其他部分
梯度累积当一个变量有多条输出路径时,梯度要求和
自动微分框架记录前向操作图,反向时自动逐节点传梯度
复杂度前向和反向都是 O(N),但反向需要额外 O(N) 空间存储中间值

📥 Code

FileViewDownload
demo.pyOpenDownload
plot_demo.pyDownload
exercise.pyOpenDownload

参考

  1. Rumelhart, D. E., Hinton, G. E., & Williams, R. J. (1986). Learning representations by back-propagating errors. Nature. [doi:10.1038/323533a0]
  2. LeCun, Y., Bottou, L., Orr, G. B., & Müller, K.-R. (1998). Efficient BackProp. Neural Networks: Tricks of the Trade. [doi:10.1007/978-3-642-35289-8_3]