链式法则 -- 复合函数求导的拆解术
链式法则是反向传播算法的数学灵魂。外导 × 内导,逐层传导。
概念解析
一元链式法则
如果 y = g(x),z = f(y),则:
\[ \frac{dz}{dx} = \frac{dz}{dy} \cdot \frac{dy}{dx} \]反向传播 = 链式法则沿计算图反向应用
x
y=g(x)
∂y/∂x
z=f(y)
∂z/∂y
L
∂L/∂z
反向传播时,梯度 从输出向输入逐层传递:每一层的「局部梯度」×「上游传下来的梯度」= 该层参数对损失的梯度。
生活例子
油门 → 油耗
油耗取决于车速,车速取决于油门深度。
「多踩一点油门,油耗增加多少?」= (油耗对车速的变化率) × (车速对油门的变化率)。
这就是链式法则:层层传导变化量,相乘得到总影响。
Python 动手实践
实例
import sympy as sp
x = sp.Symbol('x')
# f(g(x)): sin(x^2), 外层sin, 内层x^2
d_direct = sp.diff(sp.sin(x**2), x)
d_chain = sp.cos(x**2) * 2*x
print(f"直接求导: {d_direct}")
print(f"链式法则: cos(x^2)·2x = {d_chain}")
print(f"一致: {sp.simplify(d_direct-d_chain)==0}")
# 两层网络反向传播
w, xi, yi = sp.symbols('w x_i y_i')
a = sp.Symbol('a') # 中间变量单独定义为符号
a_expr = w * xi # a 关于 w 的具体表达式
loss = (a - yi)**2 # loss 定义为关于符号 a 的函数
dL_da = sp.diff(loss, a) # ∂L/∂a
da_dw = sp.diff(a_expr, w) # ∂a/∂w
dL_dw_chain = dL_da.subs(a, a_expr) * da_dw # 链式法则:代入 a=w*x_i 后相乘
loss_direct = loss.subs(a, a_expr) # 把 loss 完全展开成 w 的函数
dL_dw_direct = sp.diff(loss_direct, w)
print(f"\n链式 ∂L/∂w: {sp.simplify(dL_dw_chain)}")
print(f"直接 ∂L/∂w: {sp.simplify(dL_dw_direct)}")
print(f"一致: {sp.simplify(dL_dw_chain-dL_dw_direct)==0}")
x = sp.Symbol('x')
# f(g(x)): sin(x^2), 外层sin, 内层x^2
d_direct = sp.diff(sp.sin(x**2), x)
d_chain = sp.cos(x**2) * 2*x
print(f"直接求导: {d_direct}")
print(f"链式法则: cos(x^2)·2x = {d_chain}")
print(f"一致: {sp.simplify(d_direct-d_chain)==0}")
# 两层网络反向传播
w, xi, yi = sp.symbols('w x_i y_i')
a = sp.Symbol('a') # 中间变量单独定义为符号
a_expr = w * xi # a 关于 w 的具体表达式
loss = (a - yi)**2 # loss 定义为关于符号 a 的函数
dL_da = sp.diff(loss, a) # ∂L/∂a
da_dw = sp.diff(a_expr, w) # ∂a/∂w
dL_dw_chain = dL_da.subs(a, a_expr) * da_dw # 链式法则:代入 a=w*x_i 后相乘
loss_direct = loss.subs(a, a_expr) # 把 loss 完全展开成 w 的函数
dL_dw_direct = sp.diff(loss_direct, w)
print(f"\n链式 ∂L/∂w: {sp.simplify(dL_dw_chain)}")
print(f"直接 ∂L/∂w: {sp.simplify(dL_dw_direct)}")
print(f"一致: {sp.simplify(dL_dw_chain-dL_dw_direct)==0}")
输出:
直接求导: 2*x*cos(x**2) 链式法则: cos(x^2)·2x = 2*x*cos(x**2) 一致: True 链式 ∂L/∂w: 2*x_i*(w*x_i - y_i) 直接 ∂L/∂w: 2*x_i*(w*x_i - y_i) 一致: True
AI 中的应用场景
反向传播 = 链式法则在计算图上的系统化应用
PyTorch 的 autograd 引擎在每个张量操作时自动记录计算图。调用 loss.backward() 时,它从 loss 节点出发,沿计算图反向遍历,对每个节点应用链式法则:局部梯度 × 上游梯度 = 下游梯度。开发者不需要手写任何求导代码。
梯度检查(Gradient Check)
如果你手写了一个自定义层或损失函数,怎样验证反向传播的梯度计算正确?用数值梯度 \( (f(x+h)-f(x-h))/2h \) 作为「真值」,与 autograd 算出的梯度对比。相对误差应小于 1e-5。这是 PyTorch 中 torch.autograd.gradcheck 的原理。
计算图优化
深度学习编译器(如 TensorRT、XLA)会对计算图做优化:合并相邻操作(算子融合),减少内存读写。但必须保证链式法则下梯度计算仍然正确——这是图优化的一个重要约束。
