Create your own
Lesson illustration

用链式法则计算计算图梯度

你好。上一课我们建立了导数和偏导数的基本语言:导数描述局部变化率;对某个变量求偏导时,其余变量暂时视为常数。现在把这些局部规则组织起来,解决训练中真正的问题:一个最终的标量损失经过多步计算得到时,怎样高效地求它对早期输入或参数的梯度?

这一课的核心是链式法则与计算图。你将能够把一个表达式拆成简单节点,先做前向计算并保存中间值,再从最终输出反向逐节点计算梯度。这正是 PyTorch 的 backward()、奖励模型训练、SFT 损失反传,以及后续 PPO/DPO 梯度计算所共享的底层机制。


链式法则:把“局部变化”连接成“全局变化”

先看最简单的复合函数:

最终的 并不直接由 给出,而是先由 计算出中间变量 ,再计算 。如果 有一个很小的变化,它会先改变 ,再通过 改变

链式法则写作:

它并不是需要机械背诵的符号规则,而是两段局部变化率的组合:

  • 的变化会怎样影响中间变量
  • 的变化会怎样影响最终结果
  • 两者相乘: 的变化最终会怎样影响

例如:

于是:

因此:

这里把 代回了最终结果。直接对 求导也可以,但当模型由成百上千个算子组成时,直接展开会迅速失控。计算图的价值就在于:不必写出庞大的最终导数表达式,只需在每个局部节点计算简单导数。

下面阅读 CS231n 对加法节点和乘法节点的最小示例。它展示了“前向保存中间值,反向乘局部梯度”的标准节奏。

CS231n Deep Learning for Computer Vision

阅读 Stanford CS231n 的这一小节。它用 f(x,y,z)=(x+y)z 演示了计算图如何把复合表达式拆开,以及为什么反向计算必须从最终输出开始。

在页面开头的链式法则部分,从 这个例子 读到代码示例结束。重点跟踪前向阶段的 q 和 f 分别是多少;然后确认反向阶段中,\frac{\partial f}{\partial q} 为什么等于 z,以及 \frac{\partial f}{\partial x} 为什么还要乘上 \frac{\partial q}{\partial x}。


计算图:前向算数值,反向算梯度

计算图把一次计算表示为变量和算子的依赖结构。前向传播时,从输入出发计算每个中间变量,直到最终输出;反向传播时,从最终输出的梯度出发,按相反顺序计算每个输入和中间变量的梯度。

该计算图先计算 \(v=xy\),再计算 \(w=\ln v\);蓝色部分表示前向计算得到的变量值,绿色部分表示反向应用链式法则,得到 \(w\) 对 \(x\) 和 \(y\) 的梯度。

以图中的计算为例:

为使 有定义,需满足 。取:

前向传播:记录每个中间值

先按计算依赖顺序计算:

注意:反向传播不只需要最终的 ,还会用到 。例如,对数节点反向计算时需要知道 ,乘法节点反向计算时需要知道两个输入

这就是训练框架需要在前向阶段缓存部分激活值的基本原因。后面学习显存时,会进一步分析这种缓存为何是训练显存的重要组成。

反向传播:从输出梯度为 开始

为了表达更紧凑,定义:

其中 表示“最终输出 对变量 的梯度”。在实际自动微分实现中,这也常被称为该变量的上游梯度

由于一个标量对自身的导数为 ,反向传播的起点是:

现在反向经过对数节点。局部关系是:

因此局部导数为:

这里 ,所以:

接着反向经过乘法节点:

根据上一课的偏导数规则:

但我们真正需要的是 的梯度,因此还要乘上来自 的上游梯度:

所以最终结果是:

也可以直接验证。原式为:

固定时,对 求偏导:

代入 ,结果确实为:

直接求导与计算图反向传播的答案必须一致;区别在于,后者能系统地扩展到复杂模型。


反向传播的可执行规则

对于当前阶段的标量计算图,可以把反向传播记成一套固定流程。

  1. 按前向顺序计算所有节点值。
    记录后续求局部导数所需的输入和输出。

  2. 在最终标量输出处初始化梯度为
    若最终目标是损失 ,则:

  3. 按前向计算的相反顺序处理每个节点。
    对每个节点,写出“节点输出对节点输入”的局部导数。

  4. 局部导数乘以上游梯度。
    这一步就是链式法则。

  5. 若一个变量通过多条路径影响最终输出,累加所有路径的梯度贡献。

常见标量节点的局部梯度如下:

前向节点对第一个输入的局部导数对第二个输入的局部导数
不适用
不适用
不适用

例如,若乘法节点的输出梯度为:

且:

那么传回两个输入的梯度分别为:

这揭示一个常见现象:乘法节点会把“另一个输入”的数值带入梯度。 线性层参数梯度中出现输入激活值,就是这个局部规则在向量和矩阵上的扩展。


梯度为什么会相加:变量分叉时的多条影响路径

“沿路径相乘”只覆盖了单一路径的情形。计算图还有另一个同等重要的规则:

当一个变量在前向计算中被多个后续节点使用时,它对最终输出的总梯度等于各条路径梯度之和。

考虑:

把中间变量代回,可得:

直接求导:

下面用计算图的方式重新得到它。

前向传播中,令:

则:

反向传播从:

开始。加法节点 对两个输入的局部导数都为 ,因此:

沿 这条路径, 收到的梯度贡献是:

沿 这条路径, 收到的梯度贡献是:

由于这两条路径都从 出发并影响同一个最终损失,必须累加:

这与直接求导的结果一致:

在代码中,梯度累加通常对应于 += 的语义,而不是简单赋值。漏掉这个累加,是手写反向传播或实现自定义 loss 时很典型的错误。


一个与训练直接对应的例子:线性层、Sigmoid 与平方损失

现在把前面的规则应用到一个最小神经网络单元。设:

其中:

  • 是输入特征;
  • 是待训练参数;
  • 是线性层输出;
  • 是经过 Sigmoid 后的预测;
  • 是目标值;
  • 是单样本平方损失。

这不是语言模型常用的最终损失形式,但它清晰包含了训练中反复出现的结构:参数进入线性变换,经过非线性,再进入一个标量损失。

前向计算

取:

线性部分为:

Sigmoid 输出:

损失为:

反向计算

从损失开始:

先经过平方损失节点:

因此:

再经过 Sigmoid 节点。上一课给出了:

因此:

代入数值:

最后经过线性节点:

局部偏导为:

所以:

其中最值得记住的是参数梯度结构:

即,权重梯度等于该节点的上游梯度乘输入激活值;bias 梯度则等于上游梯度本身。后续看到线性层矩阵梯度、LoRA 参数梯度,或读训练框架中的 grad_output 时,都可以把它理解为这一标量模式的批量化版本。

下面视频把相同思想放入一个小型神经网络,重点观察“参数微调先改变线性输出,再改变激活,最后改变损失”的局部链条。

Backpropagation calculus | Deep Learning Chapter 4

观看 3Blue1Brown 的《Backpropagation calculus | Deep Learning Chapter 4》片段。它将链式法则拆成损失、激活和线性加权求和之间的三段局部导数,并解释每一项的训练含义。

观看 单个权重的梯度。留意其中损失对激活、激活对线性输出、线性输出对权重这三项如何相乘;尤其注意最后一项为何等于上一层的激活值。


手算与调试时的检查清单

在后训练代码中,你通常不会手写完整反向传播,但需要能够判断一个梯度实现是否可信。对一个简单计算图,可用以下检查方式:

  • 先明确最终标量目标。
    反向传播通常从 loss、奖励目标或某个标量 score 开始。若输出不是标量,必须先说明是在对哪个标量函数求导。

  • 前向值必须完整且可用。
    对数节点需要输入值,乘法节点需要另一侧输入,Sigmoid 节点常可利用输出 计算

  • 梯度的方向应符合直觉。
    在上面的例子中,预测 低于目标 ,所以希望增大 。由于:

    用梯度下降更新参数时会减去一个负数,从而推动 增大。

  • 同一变量重复使用时,检查是否累加梯度。
    特别是残差连接、共享参数、多个 loss 项共同训练时,同一参数往往会收到多个梯度贡献。

  • 梯度形状必须和所求变量形状一致。
    当前课以标量为主;以后进入线性层和 Transformer 时,参数梯度必须与参数张量形状一致。这是排查转置、广播和 mask 错误的重要线索。


小结

这一课建立了反向传播的最小工作模型:

  • 链式法则将复合计算分解为连续的局部导数:
  • 计算图让复杂表达式变成可管理的简单节点:前向阶段计算并缓存中间值,反向阶段按相反顺序计算梯度。

  • 对单一路径,梯度遵循“沿路径相乘”;当变量经由多条路径影响最终目标时,梯度遵循“各路径贡献相加”。

  • 对线性节点:

其参数梯度具有核心结构:

这正是神经网络参数更新能够利用激活值和误差信号的原因。

下一课将学习有限差分梯度校验:用微小扰动近似导数,并将这个数值近似与手算或自动微分结果对照。这是实现自定义 loss、奖励函数或后训练目标时非常实用的正确性检查工具。

Can't find a good explanation? Sign up and we'll make it for you

Sign up