你训练一个多层神经网络,发现怎么调学习率都不收敛。问题常常不在学习率,而在梯度——误差信号能不能从输出层一层层传回到每个权重。让这件事在数学上成立的,就是一条叫链式法则(Chain Rule)的求导公式。一句话定义:当函数是「函数套函数」时,它的导数等于「外层的导数 × 内层的导数」。
一个生活类比:流水线
想象一条流水线:原材料 x 经过第一台机器加工成半成品 u,u 再经过第二台机器加工成成品 y。如果你想知道「原材料多投入一点,成品会多产出多少」,你不必把整条流水线当一个整体去硬算,而是分两步:先算 u 对 x 的产出比,再算 y 对 u 的产出比,最后把两个比例乘起来——这就是链式法则的直觉:分步求导,再相乘。
核心公式(两层)
设
逐个符号解读:
:外层函数,把中间量 加工成最终输出 。 :内层函数,把输入 加工成中间量 。 :固定在内层这一节,问「 变一点, 跟着变多少」。 :固定在最底层,问「 变一点, 跟着变多少」。 - 两者相乘,就得到从
一路到 的总变化率。
直觉对应:流水线上两个车间的「产出比」相乘,就是头到尾的总产出比。
手算两层:拆开一个套娃
设
分步求导:
- 内层:
- 外层:
相乘:
自检:直接把
图 1:把 y=(2x+1)² 拆成内层 u=2x+1、外层 y=u²,各求导后相乘得到 8x+4。
手算三层:链子再长也一样
神经网络往往不止两层,链子会更长。设
- 最内层
: - 中间层
: - 最外层
:
三节相乘:
规律很朴素:链子有几节,导数就乘几次;每一层只管自己这一节的局部导数,链式法则负责把它们串起来。
图 2:三层函数嵌套时,每一节局部导数相乘,就得到从头到尾的总变化率。
和反向传播、梯度下降是什么关系
这才是链式法则在 AI 里真正发光的地方。一个三层神经网络,本质就是三层(甚至几十、上百层)函数嵌套:
符号逐个解读:
:损失函数,把模型预测和正确答案一比对,算出一个误差值。 :第一、二、三层各自的计算,每层都做一次「加权求和 + 激活」;三层嵌套写在一起,表示数据依次穿过——先 ,再 ,最后 。 - 分号
:把每层的两类参数隔开——分号前是输入数据(如 ),分号后是该层权重(如 )。 :最开始的输入数据。 :三层各自的权重,是网络要学的参数。
训练时,梯度下降要拿到每个权重的梯度
(这里记号从
- 输出层最近:
- 往回一层多乘一节:
- 再往回再多一节:
式中的
三件事的关系是这样的:
- 链式法则:数学工具,告诉你「套娃函数怎么求导」。
- 反向传播(Backpropagation):把链式法则高效地跑在整张网络上的算法。它从输出层开始算,每往回一层就复用「已经算好的上层梯度」,再乘一次局部导数,避免重复计算。
- 梯度下降:拿到每个权重的梯度后,真正去更新权重的那一步。
一句话收束:没有链式法则,误差就没办法从 Loss 一路传回第一层权重,多层网络根本训不动。这也是深度网络训练里常见的「梯度消失」问题的根源——链式法则里连乘的局部导数太小,乘着乘着梯度就趋近于零了。
图 3:链式法则给出求导方法,反向传播把它高效跑在网络里,梯度下降拿到梯度后更新权重。
完整代码
下面用 PyTorch 演示链式法则如何落地为自动的反向传播。对应上面两层手算例子的片段是:内层 u = 2x + 1,外层 y = u ** 2。
import torch
# —— 片段:对应「两层手算」例子 ——
x = torch.tensor(2.0, requires_grad=True) # 输入 x=2,开启梯度追踪
u = 2 * x + 1 # 内层函数 u = 2x + 1
y = u ** 2 # 外层函数 y = u^2
y.backward() # 反向传播:PyTorch 自动套链式法则
print(x.grad) # 输出 20.0,即 4(2x+1)=4×5,与手算一致
# —— 完整:一个两层全连接网络,演示梯度如何逐层传回 ——
class TinyNet(torch.nn.Module):
def __init__(self):
super().__init__()
self.fc1 = torch.nn.Linear(3, 4) # 第一层:3 维输入 → 4 维(权重 W1)
self.fc2 = torch.nn.Linear(4, 1) # 第二层:4 维 → 1 维输出(权重 W2)
def forward(self, x):
h = torch.relu(self.fc1(x)) # 中间量 u:先线性变换,再过 ReLU 激活
return self.fc2(h) # 输出 y:再一层线性变换
model = TinyNet()
x = torch.randn(1, 3) # 一条假数据
target = torch.tensor([[1.0]]) # 假目标值
y = model(x) # 前向传播:x → u → y
loss = ((y - target) ** 2).mean() # 均方误差损失 L
loss.backward() # 反向传播:沿计算图反向套链式法则,
# 自动算出 ∂L/∂W1 和 ∂L/∂W2
print(model.fc1.weight.grad) # 形状 [4, 3]:误差经「L → y → u → W1」三节链子传回小结
链式法则本身只是微积分里一条朴素的求导规则:复合函数的导数,等于各层局部导数相乘。但正是这条公式,让误差信号能顺着网络一层层倒流回每个权重,反向传播才有了数学根基,多层神经网络才训得动。
参考资料
- Chain Rule Review - Khan Academy https://www.khanacademy.org/a/chain-rule-review
- Chapter 2: How the backpropagation algorithm works - Neural Networks and Deep Learning, Michael Nielsen http://neuralnetworksanddeeplearning.com/chap2.html
- Optimization: Stochastic Gradient Descent(含链式法则与计算图) - CS231n, Stanford https://cs231n.github.io/optimization-2/