反向传播直观理解:链式法则与梯度流动
反向传播(backpropagation)是训练神经网络的核心算法。它回答了一个看似简单的问题:当网络的预测与真实答案之间存在误差时,我们应该如何调整网络中成千上万个权重,才能让下一次预测更准一些?
很多教程一上来就抛出一长串矩阵求导公式,容易让人淹没在符号里。本文换一个角度:先用直觉建立「梯度沿网络回流」的画面,再用链式法则把它变成可计算的步骤,最后用一个可以手算的最小数值示例把整条链路跑通。读完你应该能回答三件事:梯度从哪来、梯度怎么流、权重为什么这样改。
前向传播回顾
理解反向传播,先要理解前向传播。前向传播就是「数据从输入走到输出」的过程。
以一个最简单的两层全连接网络为例,假设输入是向量 x,第一层权重矩阵 W1、偏置 b1,第二层权重 W2、偏置 b2,激活函数用 tanh:
z1 = W1 · x + b1 # 第一层线性组合
a1 = tanh(z1) # 第一层非线性激活
z2 = W2 · a1 + b2 # 第二层线性组合
y_hat = z2 # 输出层(此处为线性输出)
每一层都做两件事:先对上一层输出做线性变换,再套一个非线性激活函数。非线性是神经网络能逼近任意函数的关键,也正是它让梯度回传时必须借助链式法则。
前向传播结束时,我们得到预测值 y_hat,以及一个损失值 L,衡量预测与真实标签 y 之间的差距。损失是整个网络的「成绩单」,而反向传播的目的,就是算出「每个权重对这张成绩单负多大责任」。
损失函数与梯度直觉
最常用的回归损失是均方误差(MSE):
L = 1/2 · (y_hat - y)^2
梯度衡量的是「某个量微小变化时,损失会怎么变」。对权重 w 的梯度 ∂L/∂w 表示:若把 w 轻轻拨动一点点,损失会朝哪个方向、以多大速率变化。
优化的逻辑非常朴素:如果 ∂L/∂w 为正,说明增大 w 会让损失上升,那我们就减小 w;反之就增大 w。最终更新规则就是沿负梯度走一小步:
w_new = w_old - lr · ∂L/∂w
这里的 lr 是学习率,控制每一步迈多大。问题的难点在于:网络有深有浅,输出层的权重离损失很近,可以直接求偏导;但靠近输入的权重隔着好几层,它们的梯度必须经过层层传递才能算出来。这需要一套系统的计算方法,也就是链式法则。
链式法则
链式法则是微积分里计算复合函数导数的工具。如果输出 z 依赖于 y,而 y 又依赖于 x,那么 z 对 x 的导数可以拆开:
∂z/∂x = (∂z/∂y) · (∂y/∂x)
在网络里,每一层都是对上一层的复合。损失 L 依赖于输出 z2,z2 依赖于激活 a1,a1 依赖于 z1,z1 又依赖于权重 W1。于是某个深层权重对损失的梯度,可以写成一连串局部梯度的乘积:
∂L/∂W1 = (∂L/∂z2) · (∂z2/∂a1) · (∂a1/∂z1) · (∂z1/∂W1)
反向传播的本质,就是沿着这条计算链,从损失出发一层层往回乘,把每个局部梯度就地算出来并累乘下去。它聪明的地方在于:前向时我们已经缓存了每一层的输入与输出,回传时直接复用这些值,避免重复计算,把原本指数级的代价压到了线性。
逐层回传(以两层网络为例)
把上面的直觉落到两层网络,回传顺序严格和前向相反:
第一步,从输出层往回。损失对 z2 的梯度(记做 δ2)是:
δ2 = ∂L/∂z2 = (y_hat - y)
因为输出层是线性的,这一步非常干净。
第二步,把 δ2 回传到隐藏层。它先经过权重 W2 放大或缩小,再乘上激活函数 tanh 在 z1 处的导数:
δ1 = (W2^T · δ2) ⊙ tanh'(z1)
这里 ⊙ 表示逐元素相乘,tanh’(z) = 1 - tanh(z)^2。
第三步,用每层的 δ 写出权重的梯度:
∂L/∂W2 = δ2 · a1^T
∂L/∂b2 = δ2
∂L/∂W1 = δ1 · x^T
∂L/∂b1 = δ1
关键在于 δ(有时也叫误差项、敏感度)是先算好再复用的:δ2 算一次,既用于更新 W2、b2,也用于生成 δ1,进而更新 W1、b1。这就是「反向传播」名字的由来——误差信号 δ 像水一样从输出端逆流回输入端,沿途就地生成各参数的梯度。
一个最小数值示例(Python 含前向+反向+更新)
下面用纯 Python 写一个极小的两层网络:1 个输入、2 个隐藏单元、1 个输出,激活用 tanh,损失用 MSE。我们手动把前向、反向、更新三步都写出来,不依赖任何深度学习框架,方便你对照公式一行行验证。
import math
# 超小两层网络:1 个输入、2 个隐藏单元、1 个输出
# 激活函数 tanh 及其导数
def tanh(x):
return math.tanh(x)
def tanh_grad(x):
return 1.0 - math.tanh(x) ** 2
# 初始化参数(手写固定值,便于手工追踪每一步)
w1 = [0.5, -0.3] # 输入到两个隐藏单元的权重
b1 = [0.1, -0.2]
w2 = [0.4, 0.6] # 两个隐藏单元到输出的权重
b2 = 0.0
lr = 0.1
x = 1.0
y_true = 0.8
# ---- 前向传播 ----
z1 = [w1[0] * x + b1[0], w1[1] * x + b1[1]]
a1 = [tanh(z1[0]), tanh(z1[1])]
z2 = w2[0] * a1[0] + w2[1] * a1[1] + b2
pred = z2 # 输出层不做非线性,直接当作预测值
# 均方误差损失:L = 0.5 * (pred - y_true)^2
loss = 0.5 * (pred - y_true) ** 2
# ---- 反向传播(链式法则逐层回传) ----
# 输出层对损失的梯度
dL_dz2 = (pred - y_true) # dL/dz2 = (pred - y_true)
dL_db2 = dL_dz2
# 隐藏层到输出权重的梯度
dL_dw2 = [dL_dz2 * a1[0], dL_dz2 * a1[1]]
# 误差项 δ 回传到隐藏层
dL_da1 = [dL_dz2 * w2[0], dL_dz2 * w2[1]]
dL_dz1 = [dL_da1[0] * tanh_grad(z1[0]), dL_da1[1] * tanh_grad(z1[1])]
dL_dw1 = [dL_dz1[0] * x, dL_dz1[1] * x]
dL_db1 = [dL_dz1[0], dL_dz1[1]]
# ---- 参数更新(SGD) ----
w2[0] -= lr * dL_dw2[0]
w2[1] -= lr * dL_dw2[1]
b2 -= lr * dL_db2
w1[0] -= lr * dL_dw1[0]
w1[1] -= lr * dL_dw1[1]
b1[0] -= lr * dL_db1[0]
b1[1] -= lr * dL_db1[1]
print("loss =", loss)
print("更新后 w2 =", w2, "b2 =", b2)
这段代码里没有魔法:dL_dz2 就是上一节的 δ2,dL_dz1 就是 δ1,所有权重梯度都严格按「δ 乘上对应层的输入」得到。把这几行跑一遍,你会看到损失从某个正数开始下降,权重朝「让预测更接近 0.8」的方向移动。这正是反向传播在做的事——把全局误差拆解成每个参数的局部责任,再据此更新。
常见误区(梯度消失/爆炸)
理解了原理,也要警惕两个经典陷阱。
梯度消失。当网络很深、激活函数处在导数很小的区间(比如 tanh 两端、或早期常用的 sigmoid 饱和区)时,链式法则里一连串小于 1 的因子相乘,会让靠近输入层的梯度指数级缩小。结果是浅层权重几乎学不到东西,深层网络难以训练。缓解手段包括换用 ReLU 类激活函数、合理的权重初始化(如 Xavier、He)、以及残差连接(ResNet)等。
梯度爆炸。如果局部梯度普遍大于 1,连乘之后又会指数级放大,导致更新步长巨大、损失剧烈震荡甚至 NaN。这在循环神经网络里尤其常见。常用对策是梯度裁剪(gradient clipping),即给梯度范数设上限,以及在时序结构上做归一化。
值得强调的是:反向传播本身只是「高效求梯度」的算法,训练稳不稳,取决于梯度在深层链路里是被逐渐抹平还是被不断放大。这也是为什么现代架构在激活、初始化、归一化上花了很多心思。
小结
- 反向传播是沿计算图从损失出发、用链式法则逐层回传梯度,从而算出每个参数对损失的贡献。
- 它的效率来自「复用前向缓存的中间结果」,把本应指数级的求导代价降为线性。
- 实践中梯度以误差项 δ 的形式逆流:δ2 直接可得,δ1 由 δ2 经权重与激活导数生成,权重梯度由 δ 乘对应层输入得到。
- 深层网络的梯度消失与爆炸,是设计激活函数、初始化与归一化策略时真正要解决的问题。
- 想彻底吃透,最好的办法是抛开框架手写一遍前向与反向,就像上文的最小示例那样。
参考与延伸阅读
- Rumelhart D. E., Hinton G. E., Williams R. J. 1986. Learning representations by back-propagating errors. Nature, volume 323, pages 533–536.(反向传播原始论文,Nature 官网可查,DOI 10.1038/323533a0)
- Goodfellow I., Bengio Y., Courville A. 2016. Deep Learning. MIT Press, 第 6.5 节 Back-Propagation and Other Differentiation Algorithms.(系统推导反向传播与计算图,官方免费版见 deeplearningbook.org)
- 3Blue1Brown. Deep Learning 系列视频,第 3 集 Backpropagation, intuitively 与第 4 集 Backpropagation calculus(YouTube,直观与微积分两种视角讲解反向传播)。
- Karpathy A. micrograd:一个极简的标量自动求导引擎与神经网络库,地址 github.com/karpathy/micrograd(约百行代码实现反向传播,适合精读源码)。
- Rumelhart D. E., McClelland J. L. (编) 1986. Parallel Distributed Processing: Explorations in the Microstructure of Cognition, Vol. 1: Foundations. MIT Press, pages 318–362.(反向传播更完整的早期阐述,亦被 Nature 1986 论文列为参考文献)。