链式法则是什么

链式法则封面

你训练一个多层神经网络,发现怎么调学习率都不收敛。问题常常不在学习率,而在梯度——误差信号能不能从输出层一层层传回到每个权重。让这件事在数学上成立的,就是一条叫链式法则(Chain Rule)的求导公式。一句话定义:当函数是「函数套函数」时,它的导数等于「外层的导数 × 内层的导数」

一个生活类比:流水线

想象一条流水线:原材料 x 经过第一台机器加工成半成品 uu 再经过第二台机器加工成成品 y。如果你想知道「原材料多投入一点,成品会多产出多少」,你不必把整条流水线当一个整体去硬算,而是分两步:先算 ux 的产出比,再算 yu 的产出比,最后把两个比例乘起来——这就是链式法则的直觉:分步求导,再相乘

核心公式(两层)

u=g(x)u = g(x)y=f(u)y = f(u),则 yyxx 的导数是:

dydx=dydududx\frac{dy}{dx} = \frac{dy}{du} \cdot \frac{du}{dx}

逐个符号解读:

  • y=f(u)y = f(u):外层函数,把中间量 uu 加工成最终输出 yy
  • u=g(x)u = g(x):内层函数,把输入 xx 加工成中间量 uu
  • dydu\dfrac{dy}{du}:固定在内层这一节,问「uu 变一点,yy 跟着变多少」。
  • dudx\dfrac{du}{dx}:固定在最底层,问「xx 变一点,uu 跟着变多少」。
  • 两者相乘,就得到从 xx 一路到 yy 的总变化率。

直觉对应:流水线上两个车间的「产出比」相乘,就是头到尾的总产出比。

手算两层:拆开一个套娃

y=(2x+1)2y = (2x+1)^2。它看着是一个式子,其实是两层套娃:内层 u=2x+1u = 2x+1,外层 y=u2y = u^2

分步求导:

  • 内层:dudx=2\dfrac{du}{dx} = 2
  • 外层:dydu=2u=2(2x+1)\dfrac{dy}{du} = 2u = 2(2x+1)

相乘:

dydx=2(2x+1)2=4(2x+1)=8x+4\frac{dy}{dx} = 2(2x+1) \cdot 2 = 4(2x+1) = 8x + 4

自检:直接把 yy 展开成 4x2+4x+14x^2 + 4x + 1,求导得 8x+48x + 4,两者一致,结果可信。

两层套娃拆解

图 1:把 y=(2x+1)² 拆成内层 u=2x+1、外层 y=u²,各求导后相乘得到 8x+4。

手算三层:链子再长也一样

神经网络往往不止两层,链子会更长。设 y=σ(3(2x)+1)y = \sigma(3(2x)+1),其中 σ\sigma 是 sigmoid 激活函数。拆成三节:

  • 最内层 a=2xa = 2xdadx=2\dfrac{da}{dx} = 2
  • 中间层 u=3a+1u = 3a + 1duda=3\dfrac{du}{da} = 3
  • 最外层 y=σ(u)y = \sigma(u)dydu=σ(u)(1σ(u))\dfrac{dy}{du} = \sigma(u)\bigl(1-\sigma(u)\bigr)

三节相乘:

dydx=dydududadadx=σ(u)(1σ(u))32=6σ(u)(1σ(u))\frac{dy}{dx} = \frac{dy}{du} \cdot \frac{du}{da} \cdot \frac{da}{dx} = \sigma(u)\bigl(1-\sigma(u)\bigr) \cdot 3 \cdot 2 = 6\,\sigma(u)\bigl(1-\sigma(u)\bigr)

规律很朴素:链子有几节,导数就乘几次;每一层只管自己这一节的局部导数,链式法则负责把它们串起来

三层链式求导

图 2:三层函数嵌套时,每一节局部导数相乘,就得到从头到尾的总变化率。

和反向传播、梯度下降是什么关系

这才是链式法则在 AI 里真正发光的地方。一个三层神经网络,本质就是三层(甚至几十、上百层)函数嵌套:

L=Loss(f3(f2(f1(x;W1);W2);W3))L = \text{Loss}\bigl(f_3(f_2(f_1(x;\,W_1);\,W_2);\,W_3)\bigr)

符号逐个解读:

  • Loss\text{Loss}:损失函数,把模型预测和正确答案一比对,算出一个误差值。
  • f1,f2,f3f_1, f_2, f_3:第一、二、三层各自的计算,每层都做一次「加权求和 + 激活」;三层嵌套写在一起,表示数据依次穿过——先 f1f_1,再 f2f_2,最后 f3f_3
  • 分号 ;;:把每层的两类参数隔开——分号前是输入数据(如 xx),分号后是该层权重(如 W1W_1)。
  • xx:最开始的输入数据。
  • W1,W2,W3W_1, W_2, W_3:三层各自的权重,是网络要学的参数。

训练时,梯度下降要拿到每个权重的梯度 LW1\dfrac{\partial L}{\partial W_1}LW2\dfrac{\partial L}{\partial W_2}LW3\dfrac{\partial L}{\partial W_3} 才能更新权重。怎么算?就是反复套链式法则:

(这里记号从 dd 换成了 \partial\partial 是偏导符号,多变量时用,和前面的 dd 同是求导记号,含义一致。)

  • 输出层最近:LW3=LyyW3\dfrac{\partial L}{\partial W_3} = \dfrac{\partial L}{\partial y} \cdot \dfrac{\partial y}{\partial W_3}
  • 往回一层多乘一节:LW2=Lyyu2u2W2\dfrac{\partial L}{\partial W_2} = \dfrac{\partial L}{\partial y} \cdot \dfrac{\partial y}{\partial u_2} \cdot \dfrac{\partial u_2}{\partial W_2}
  • 再往回再多一节:LW1=Lyyu2u2u1u1W1\dfrac{\partial L}{\partial W_1} = \dfrac{\partial L}{\partial y} \cdot \dfrac{\partial y}{\partial u_2} \cdot \dfrac{\partial u_2}{\partial u_1} \cdot \dfrac{\partial u_1}{\partial W_1}

式中的 u1,u2u_1, u_2 是中间隐藏层的中转量:u2u_2 是第二层的输出、u1u_1 是第一层的输出——数据从 xx 出发先经 f1f_1 变成 u1u_1,再经 f2f_2 变成 u2u_2,最后才到达最终输出 yy

三件事的关系是这样的:

  • 链式法则:数学工具,告诉你「套娃函数怎么求导」。
  • 反向传播(Backpropagation):把链式法则高效地跑在整张网络上的算法。它从输出层开始算,每往回一层就复用「已经算好的上层梯度」,再乘一次局部导数,避免重复计算。
  • 梯度下降:拿到每个权重的梯度后,真正去更新权重的那一步。

一句话收束:没有链式法则,误差就没办法从 Loss 一路传回第一层权重,多层网络根本训不动。这也是深度网络训练里常见的「梯度消失」问题的根源——链式法则里连乘的局部导数太小,乘着乘着梯度就趋近于零了。

三者分工关系

图 3:链式法则给出求导方法,反向传播把它高效跑在网络里,梯度下降拿到梯度后更新权重。

完整代码

下面用 PyTorch 演示链式法则如何落地为自动的反向传播。对应上面两层手算例子的片段是:内层 u = 2x + 1,外层 y = u ** 2

python
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」三节链子传回

小结

链式法则本身只是微积分里一条朴素的求导规则:复合函数的导数,等于各层局部导数相乘。但正是这条公式,让误差信号能顺着网络一层层倒流回每个权重,反向传播才有了数学根基,多层神经网络才训得动。

参考资料

  1. Chain Rule Review - Khan Academy https://www.khanacademy.org/a/chain-rule-review
  2. Chapter 2: How the backpropagation algorithm works - Neural Networks and Deep Learning, Michael Nielsen http://neuralnetworksanddeeplearning.com/chap2.html
  3. Optimization: Stochastic Gradient Descent(含链式法则与计算图) - CS231n, Stanford https://cs231n.github.io/optimization-2/