你好。上一课我们建立了导数和偏导数的基本语言:导数描述局部变化率;对某个变量求偏导时,其余变量暂时视为常数。现在把这些局部规则组织起来,解决训练中真正的问题:一个最终的标量损失经过多步计算得到时,怎样高效地求它对早期输入或参数的梯度?
这一课的核心是链式法则与计算图。你将能够把一个表达式拆成简单节点,先做前向计算并保存中间值,再从最终输出反向逐节点计算梯度。这正是 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}。
计算图:前向算数值,反向算梯度
计算图把一次计算表示为变量和算子的依赖结构。前向传播时,从输入出发计算每个中间变量,直到最终输出;反向传播时,从最终输出的梯度出发,按相反顺序计算每个输入和中间变量的梯度。

以图中的计算为例:
为使 有定义,需满足 。取:
前向传播:记录每个中间值
先按计算依赖顺序计算:
注意:反向传播不只需要最终的 ,还会用到 、 和 。例如,对数节点反向计算时需要知道 ,乘法节点反向计算时需要知道两个输入 。
这就是训练框架需要在前向阶段缓存部分激活值的基本原因。后面学习显存时,会进一步分析这种缓存为何是训练显存的重要组成。
反向传播:从输出梯度为 开始
为了表达更紧凑,定义:
其中 表示“最终输出 对变量 的梯度”。在实际自动微分实现中,这也常被称为该变量的上游梯度。
由于一个标量对自身的导数为 ,反向传播的起点是:
现在反向经过对数节点。局部关系是:
因此局部导数为:
这里 ,所以:
接着反向经过乘法节点:
根据上一课的偏导数规则:
但我们真正需要的是 对 的梯度,因此还要乘上来自 的上游梯度:
所以最终结果是:
也可以直接验证。原式为:
在 固定时,对 求偏导:
代入 ,结果确实为:
直接求导与计算图反向传播的答案必须一致;区别在于,后者能系统地扩展到复杂模型。
反向传播的可执行规则
对于当前阶段的标量计算图,可以把反向传播记成一套固定流程。
-
按前向顺序计算所有节点值。
记录后续求局部导数所需的输入和输出。 -
在最终标量输出处初始化梯度为 。
若最终目标是损失 ,则: -
按前向计算的相反顺序处理每个节点。
对每个节点,写出“节点输出对节点输入”的局部导数。 -
局部导数乘以上游梯度。
这一步就是链式法则。 -
若一个变量通过多条路径影响最终输出,累加所有路径的梯度贡献。
常见标量节点的局部梯度如下:
| 前向节点 | 对第一个输入的局部导数 | 对第二个输入的局部导数 |
|---|---|---|
| 不适用 | ||
| 不适用 | ||
| 不适用 |
例如,若乘法节点的输出梯度为:
且:
那么传回两个输入的梯度分别为:
这揭示一个常见现象:乘法节点会把“另一个输入”的数值带入梯度。 线性层参数梯度中出现输入激活值,就是这个局部规则在向量和矩阵上的扩展。
梯度为什么会相加:变量分叉时的多条影响路径
“沿路径相乘”只覆盖了单一路径的情形。计算图还有另一个同等重要的规则:
当一个变量在前向计算中被多个后续节点使用时,它对最终输出的总梯度等于各条路径梯度之和。
考虑:
把中间变量代回,可得:
直接求导:
下面用计算图的方式重新得到它。
前向传播中,令:
则:
反向传播从:
开始。加法节点 对两个输入的局部导数都为 ,因此:
沿 这条路径, 收到的梯度贡献是:
沿 这条路径, 收到的梯度贡献是:
由于这两条路径都从 出发并影响同一个最终损失,必须累加:
这与直接求导的结果一致:
在代码中,梯度累加通常对应于 += 的语义,而不是简单赋值。漏掉这个累加,是手写反向传播或实现自定义 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