链式法则与计算图反向传播 Chain Rule & Backpropagation
复合函数求导不是死记硬背的算术符号,而是放大倍数的层层连乘;当变量沿多条路径汇聚时,各路影响相加即为全导数。现代深度学习的“反向传播”,正是链式法则在拓扑计算图上的高效逆向逆行。
复合放大与反向求导 · 交互实验室
拖动转角滑块观察:每个齿轮将前一级的微小位移按齿轮半径比放大或缩小。两级连续放大,总变化率等于各级放大倍数的连乘:dy/dx = (dy/du) · (du/dx)。如果多级传动比小于 1,输出几乎纹丝不动(梯度消失直觉)。
当输入 t 同时沿着支路 u 和支路 v 影响终端 z 时,两条路径的影响在终点合并汇聚。全微分链式法则表明:总导数 dz/dt 是所有分支导数之和。反向传播时,下游回传至同一个变量的多个分支梯度必须进行相加累积(+=)。
正向计算自左向右评估节点数值;反向传播自右向左,根据拓扑排序逆向点亮伴随量 v̄ᵢ = ∂L/∂vᵢ。注意观察预设 2 中 y 节点接收两条支路累加,以及预设 3 中跳跃连接为 x 带来的恒定 +1.0 梯度穿透。
在深度学习训练中,输入权重规模动辄数千万到数千亿(N ≫ 1),而损失函数恒为单一标量(M = 1)。反向模式自动微分的单次计算复杂度完全与参数量 N 解耦,这是让超大规模神经网络训练成为现实的核心算法突破。
极简数学原理:从齿轮乘法到拓扑伴随量
1. 单变量复合求导:放大倍数的连乘
设 $y = f(u)$ 且 $u = g(x)$。当自变量 $x$ 扰动微元 $dx$ 时,中间变量 $u$ 的响应为 $du \approx g'(x)dx$;中间变量 $u$ 的变化又引起输出 $y$ 的响应 $dy \approx f'(u)du$。将其代入即得:
$$\frac{dy}{dx} = \frac{dy}{du} \cdot \frac{du}{dx} = f'(g(x)) \cdot g'(x)$$
核心直觉:变化率在传递时如同齿轮变速,前一级的放大倍数乘以第二级的放大倍数。若有多层复合 $y = f_L(\dots f_1(x))$,总导数就是各层局部导数的长串连乘。
2. 多元链式法则:分支路径相加
若变量 $t$ 通过多条不同的中间变量路径(例如 $u(t)$ 和 $v(t)$)同时作用于最终函数 $z = f(u, v)$,全微分公式表明总变化率是所有路径贡献的累加:
$$\frac{dz}{dt} = \frac{\partial z}{\partial u}\frac{du}{dt} + \frac{\partial z}{\partial v}\frac{dv}{dt}$$
口诀:「沿路径连乘,跨分支相加」。在深度学习计算图中,任何被多处使用的节点(如权重共享或残差跳连),其梯度必定是来自各下游分支梯度的代数和。
3. 计算图伴随量与反向模式 AD
在计算图中,定义每个节点 $v_i$ 的伴随量(Adjoint)为最终标量目标 $L$ 对该节点的偏导数:
$$\bar{v}_i \triangleq \frac{\partial L}{\partial v_i} = \sum_{j \in \text{children}(i)} \bar{v}_j \frac{\partial v_j}{\partial v_i}$$
从输出端开始初始化 $\bar{L} = \frac{\partial L}{\partial L} = 1.0$。按拓扑排序的逆序依次处理每个节点,将已计算出的下游伴随量 $\bar{v}_j$ 与局部偏导数相乘并累加给父节点,直到回传至所有输入参数。
4. 高维雅可比与向量-雅可比乘积 (VJP)
当输入与输出均为向量时,一阶导数为雅可比矩阵 $J \in \mathbb{R}^{m \times n}$,其中 $J_{ij} = \frac{\partial y_i}{\partial x_j}$。反向模式自动微分的数学本质,正是从后向前执行向量-雅可比乘积(Vector-Jacobian Product, VJP):
$$v^T J = \left[ \bar{y}_1, \dots, \bar{y}_m \right] \begin{bmatrix} \frac{\partial y_1}{\partial x_1} & \dots & \frac{\partial y_1}{\partial x_n} \\ \vdots & \ddots & \vdots \\ \frac{\partial y_m}{\partial x_1} & \dots & \frac{\partial y_m}{\partial x_n} \end{bmatrix}$$
反向模式从来不显式构造巨大的 $m \times n$ 雅可比矩阵,而是直接计算行向量与局部算子的矩阵乘积,大幅节约显存与计算时间。
为什么学这步?
整个深度学习大厦(卷积神经网络、Transformer 大语言模型、强化学习策略梯度)都建立在“通过损失函数梯度调整数以亿计参数”这一基础机制上。如果采用数值差分,每个参数都要前传一次,训练一次 GPT 级别的大模型需要数万年;而反向传播依靠链式法则与伴随量反向拓扑流,只需一次前传和一次反传即可瞬时提取所有参数的精确导数。真正搞懂链式法则与计算图,你就掌握了 PyTorch / JAX 自动微分引擎的核心心跳。
发生了什么
- 复合函数求导本质是微观局部放大倍数的串联相乘
- 多路径汇聚时遵循全微分分支相加原理(`grad += incoming`)
- 反向模式 AD 以 O(1) 复杂度求出标量损失对所有参数的梯度
常见陷阱 ⚠️
- 切忌混淆乘法与加法:复合是连乘,多分支汇聚才是相加
- 外层导数 f'(u) 必须在中间值 u=g(x) 处计算,绝不能直接代入 x
- 深层多层连续相乘若模长小于 1 会引发严重的梯度消失
本章小结
- 单变量链式法则:dy/dx = (dy/du) · (du/dx)
- 计算图伴随量递推:v̄ᵢ = Σ v̄_j · (∂v_j/∂v_i)
- ResNet 跳跃连接通过 ∂(x+F)/∂x = I + ∂F/∂x 恒保梯度穿透
🛣️ ResNet 为什么能训练上千层?梯度高速公路的数学证明
在 ResNet 诞生之前,超过 20 层的网络往往无法训练,因为深层梯度的传递是数十个雅可比矩阵的连乘:$\frac{\partial L}{\partial x_1} = \frac{\partial L}{\partial x_L} \prod_{l=1}^{L-1} \frac{\partial x_{l+1}}{\partial x_l}$。只要各层的特征值模长略小于 1(例如 0.8),经过 30 层后 $0.8^{30} \approx 0.0012$,浅层梯度几乎彻底归零。
何恺明等人引入残差连接 $x_{l+1} = x_l + F(x_l)$,对输入 $x_l$ 求导得到:
$$\frac{\partial x_{l+1}}{\partial x_l} = I + \frac{\partial F(x_l)}{\partial x_l}$$
将多层连乘展开后,表达式中天然包含一项纯由恒等矩阵构成的乘积:$I \cdot I \dots I = I$!这意味着:即使所有权重层分支的导数 $\frac{\partial F}{\partial x}$ 全部变为零甚至死掉,梯度仍然能沿着跳跃连接以大小为 1 的无损幅度直达最前层的输入参数,从而彻底根治了梯度消失。
📚 参考文献与延伸阅读
- Baydin et al. (2018), Automatic Differentiation in Machine Learning: A Survey, JMLR — 计算机科学与应用数学界对正向/反向 AD 的权威综述
- Goodfellow, Bengio & Courville, Deep Learning, MIT Press, Ch. 6.5「Back-Propagation and Other Differentiation Algorithms」 — 深度学习圣经关于计算图反向传播的标准推导
- He et al. (2016), Deep Residual Learning for Image Recognition — ResNet 残差跳连经典论文,对照阅读 → DL 第 9 步 · ResNet
- 对照阅读 → ML 第 5 步 · 多层感知机与反向传播、DL 第 1 步 · 梯度消失分析、DL 第 8 步 · RNN 沿时间反向传播 (BPTT)
📝 课后练习
检验你的理解——答对为止