MATH ML Learning Lab
06 / 09
第 6 步 · 复合求导与计算图反向传播

链式法则与计算图反向传播 Chain Rule & Backpropagation

复合函数求导不是死记硬背的算术符号,而是放大倍数的层层连乘;当变量沿多条路径汇聚时,各路影响相加即为全导数。现代深度学习的“反向传播”,正是链式法则在拓扑计算图上的高效逆向逆行。

阅读+实操 25 min 交互实验 4 个 难度 中阶

复合放大与反向求导 · 交互实验室

① 齿轮联动装置:单变量复合求导与放大倍数连乘
第一级放大 du/dx: 1.50 第二级放大 dy/du: 2.00 总放大倍数 dy/dx: 3.00

拖动转角滑块观察:每个齿轮将前一级的微小位移按齿轮半径比放大或缩小。两级连续放大,总变化率等于各级放大倍数的连乘:dy/dx = (dy/du) · (du/dx)。如果多级传动比小于 1,输出几乎纹丝不动(梯度消失直觉)。

② 分支汇流池:多元链式法则与分支路径相加
支路 A 贡献 (∂z/∂u)(du/dt): +1.80 支路 B 贡献 (∂z/∂v)(dv/dt): +1.20 总全导数 dz/dt: +3.00

当输入 t 同时沿着支路 u 和支路 v 影响终端 z 时,两条路径的影响在终点合并汇聚。全微分链式法则表明:总导数 dz/dt 是所有分支导数之和。反向传播时,下游回传至同一个变量的多个分支梯度必须进行相加累积(+=)。

③ 交互式计算图求导器:拓扑伴随量(Adjoint)逆向回传
点击「正向求值」或「单步反向回传」观察计算图节点的动态点亮与伴随量累加

正向计算自左向右评估节点数值;反向传播自右向左,根据拓扑排序逆向点亮伴随量 v̄ᵢ = ∂L/∂vᵢ。注意观察预设 2 中 y 节点接收两条支路累加,以及预设 3 中跳跃连接为 x 带来的恒定 +1.0 梯度穿透。

④ 前向模式 vs 反向模式自动微分:为什么深度学习离不开反向传播?
输入参数量 N (权重):
10,000
输出维数 M (目标损失):
1 (标量损失)
前向模式所需遍数:
10,000 遍
反向模式所需遍数:
1 遍
输入 N = 10,000,输出 M = 1 时,反向模式只需 1 遍逆向回传 即可提取全梯度,前向模式需要 10,000 遍 前传,提速比达到 10,000 倍

在深度学习训练中,输入权重规模动辄数千万到数千亿(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 的无损幅度直达最前层的输入参数,从而彻底根治了梯度消失。

📚 参考文献与延伸阅读

📝 课后练习

检验你的理解——答对为止