Lesson illustration

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

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

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


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

先看最简单的复合函数:

u=g(x)u=g(x) L=f(u)L=f(u)

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

链式法则写作:

dLdx=dLdududx\frac{dL}{dx} = \frac{dL}{du} \frac{du}{dx}

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

  • dudx\frac{du}{dx}xx 的变化会怎样影响中间变量 uu
  • dLdu\frac{dL}{du}uu 的变化会怎样影响最终结果 LL
  • 两者相乘:xx 的变化最终会怎样影响 LL

例如:

u=x2u=x^2 L=lnuL=\ln u

于是:

dLdu=1u\frac{dL}{du}=\frac{1}{u} dudx=2x\frac{du}{dx}=2x

因此:

dLdx=1u2x=2xx2=2x\frac{dL}{dx} = \frac{1}{u}\cdot 2x = \frac{2x}{x^2} = \frac{2}{x}

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

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

{"type":"reading","par_intro":"阅读 Stanford CS231n 的这一小节。它用 \\(f(x,y,z)=(x+y)z\\) 演示了计算图如何把复合表达式拆开,以及为什么反向计算必须从最终输出开始。","par_directions":"在页面开头的链式法则部分,从 <span data-type=\"resource_reading_textrange\" data-resource-subitem-id=\"c47e6f52\" data-range-start=\"Lets now start to consider more complicated expressions\" data-range-end=\"This is the simplest example of backpropagation.\">这个例子</span> 读到代码示例结束。重点跟踪前向阶段的 \\(q\\) 和 \\(f\\) 分别是多少;然后确认反向阶段中,\\(\\frac{\\partial f}{\\partial q}\\) 为什么等于 \\(z\\),以及 \\(\\frac{\\partial f}{\\partial x}\\) 为什么还要乘上 \\(\\frac{\\partial q}{\\partial x}\\)。","learning_duration":"8 minutes","url":"https://cs231n.github.io/optimization-2","title":"CS231n Deep Learning for Computer Vision","isV2":true,"blockId":"60addbf8-c618-450e-b420-e1ccf0e7e8eb","lessonId":"a36a726f-6416-4956-a06e-4dc48e3f0eab"}




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

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

{"type":"image","url":"https://pytorch.org/wp-content/uploads/2021/06/extended_computational_graph.png","caption":"该计算图先计算 \\(v=xy\\),再计算 \\(w=\\ln v\\);蓝色部分表示前向计算得到的变量值,绿色部分表示反向应用链式法则,得到 \\(w\\) 对 \\(x\\) 和 \\(y\\) 的梯度。","isV2":true,"blockId":"f75e7eaa-d977-4325-9622-f2e3d0fb125a","lessonId":"a36a726f-6416-4956-a06e-4dc48e3f0eab"}



以图中的计算为例:

v=xyv=xy w=lnvw=\ln v

为使 lnv\ln v 有定义,需满足 v>0v>0。取:

x=2,y=3x=2,\qquad y=3

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

先按计算依赖顺序计算:

v=xy=2×3=6v=xy=2\times 3=6 w=lnv=ln6w=\ln v=\ln 6

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

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

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

为了表达更紧凑,定义:

uˉ=wu\bar{u} = \frac{\partial w}{\partial u}

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

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

wˉ=ww=1\bar{w} = \frac{\partial w}{\partial w} = 1

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

w=lnvw=\ln v

因此局部导数为:

wv=1v\frac{\partial w}{\partial v} = \frac{1}{v}

这里 v=6v=6,所以:

vˉ=wv=16\bar{v} = \frac{\partial w}{\partial v} = \frac{1}{6}

接着反向经过乘法节点:

v=xyv=xy

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

vx=y\frac{\partial v}{\partial x}=y vy=x\frac{\partial v}{\partial y}=x

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

xˉ=wx=wvvx=16×3=12\bar{x} = \frac{\partial w}{\partial x} = \frac{\partial w}{\partial v} \frac{\partial v}{\partial x} = \frac{1}{6}\times 3 = \frac{1}{2} yˉ=wy=wvvy=16×2=13\bar{y} = \frac{\partial w}{\partial y} = \frac{\partial w}{\partial v} \frac{\partial v}{\partial y} = \frac{1}{6}\times 2 = \frac{1}{3}

所以最终结果是:

wx=12\frac{\partial w}{\partial x}=\frac{1}{2} wy=13\frac{\partial w}{\partial y}=\frac{1}{3}

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

w=ln(xy)w=\ln(xy)

yy 固定时,对 xx 求偏导:

wx=1xyy=1x\frac{\partial w}{\partial x} = \frac{1}{xy}\cdot y = \frac{1}{x}

代入 x=2x=2,结果确实为:

wx=12\frac{\partial w}{\partial x} = \frac{1}{2}

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

{
  "type": "exercise",
  "id": "ec139b55-980f-4f9d-9ca1-8a59347f1db1"
}

反向传播的可执行规则

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

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

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

    LL=1\frac{\partial L}{\partial L}=1
  3. 按前向计算的相反顺序处理每个节点。
    对每个节点,写出“节点输出对节点输入”的局部导数。

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

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

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

前向节点对第一个输入的局部导数对第二个输入的局部导数
q=a+bq=a+bqa=1\frac{\partial q}{\partial a}=1qb=1\frac{\partial q}{\partial b}=1
q=abq=abqa=b\frac{\partial q}{\partial a}=bqb=a\frac{\partial q}{\partial b}=a
q=a2q=a^2qa=2a\frac{\partial q}{\partial a}=2a不适用
q=lnaq=\ln aqa=1a\frac{\partial q}{\partial a}=\frac{1}{a}不适用
q=eaq=e^aqa=ea=q\frac{\partial q}{\partial a}=e^a=q不适用

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

qˉ=Lq\bar{q}=\frac{\partial L}{\partial q}

且:

q=abq=ab

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

aˉ=La=qˉb\bar{a} = \frac{\partial L}{\partial a} = \bar{q}\cdot b bˉ=Lb=qˉa\bar{b} = \frac{\partial L}{\partial b} = \bar{q}\cdot a

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


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

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

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

考虑:

p=x2p=x^2 q=3xq=3x L=p+qL=p+q

把中间变量代回,可得:

L=x2+3xL=x^2+3x

直接求导:

dLdx=2x+3\frac{dL}{dx}=2x+3

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

前向传播中,令:

x=2x=2

则:

p=4p=4 q=6q=6 L=10L=10

反向传播从:

Lˉ=1\bar{L}=1

开始。加法节点 L=p+qL=p+q 对两个输入的局部导数都为 11,因此:

pˉ=1\bar{p}=1 qˉ=1\bar{q}=1

沿 p=x2p=x^2 这条路径,xx 收到的梯度贡献是:

xˉp=pˉ2x=1×4=4\bar{x}_{p} = \bar{p}\cdot 2x = 1\times 4 = 4

沿 q=3xq=3x 这条路径,xx 收到的梯度贡献是:

xˉq=qˉ3=1×3=3\bar{x}_{q} = \bar{q}\cdot 3 = 1\times 3 = 3

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

xˉ=xˉp+xˉq=4+3=7\bar{x} = \bar{x}_{p} + \bar{x}_{q} = 4+3 = 7

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

2x+3=2×2+3=72x+3=2\times 2+3=7

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

{
  "type": "exercise",
  "id": "62aa6ef4-4f6f-4f1f-b65a-8dbdca14d043"
}

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

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

z=wx+bz=wx+b a=σ(z)a=\sigma(z) L=12(ay)2L=\frac{1}{2}(a-y)^2

其中:

  • xx 是输入特征;
  • w,bw,b 是待训练参数;
  • zz 是线性层输出;
  • aa 是经过 Sigmoid 后的预测;
  • yy 是目标值;
  • LL 是单样本平方损失。

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

前向计算

取:

x=2,w=0.5,b=1,y=1x=2,\qquad w=0.5,\qquad b=-1,\qquad y=1

线性部分为:

z=wx+b=0.5×21=0z=wx+b=0.5\times 2-1=0

Sigmoid 输出:

a=σ(0)=0.5a=\sigma(0)=0.5

损失为:

L=12(0.51)2=0.125L=\frac{1}{2}(0.5-1)^2=0.125

反向计算

从损失开始:

Lˉ=1\bar{L}=1

先经过平方损失节点:

L=12(ay)2L=\frac{1}{2}(a-y)^2

因此:

aˉ=La=ay=0.51=0.5\bar{a} = \frac{\partial L}{\partial a} = a-y = 0.5-1 = -0.5

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

az=a(1a)\frac{\partial a}{\partial z} = a(1-a)

因此:

zˉ=Lz=aˉa(1a)\bar{z} = \frac{\partial L}{\partial z} = \bar{a}\cdot a(1-a)

代入数值:

zˉ=0.5×0.5×0.5=0.125\bar{z} = -0.5\times 0.5\times 0.5 = -0.125

最后经过线性节点:

z=wx+bz=wx+b

局部偏导为:

zw=x\frac{\partial z}{\partial w}=x zb=1\frac{\partial z}{\partial b}=1 zx=w\frac{\partial z}{\partial x}=w

所以:

Lw=zˉx=0.125×2=0.25\frac{\partial L}{\partial w} = \bar{z}\cdot x = -0.125\times 2 = -0.25 Lb=zˉ1=0.125\frac{\partial L}{\partial b} = \bar{z}\cdot 1 = -0.125 Lx=zˉw=0.125×0.5=0.0625\frac{\partial L}{\partial x} = \bar{z}\cdot w = -0.125\times 0.5 = -0.0625

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

Lw=Lzx\frac{\partial L}{\partial w} = \frac{\partial L}{\partial z}\cdot x Lb=Lz\frac{\partial L}{\partial b} = \frac{\partial L}{\partial z}

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

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

{"type":"video","title":"Backpropagation calculus | Deep Learning Chapter 4","learning_duration":151,"video_id":"tIeHLnjs5U8","par_intro":"观看 3Blue1Brown 的《Backpropagation calculus | Deep Learning Chapter 4》片段。它将链式法则拆成损失、激活和线性加权求和之间的三段局部导数,并解释每一项的训练含义。","par_directions":"观看 <span data-type=\"resource_video_timerange\" data-resource-subitem-id=\"7b668502\" data-range-start=\"161\" data-range-end=\"312\">单个权重的梯度</span>。留意其中损失对激活、激活对线性输出、线性输出对权重这三项如何相乘;尤其注意最后一项为何等于上一层的激活值。","video_duration":618,"isV2":true,"blockId":"a6bbfa1a-a6ec-4007-9e94-2a02a689e81e","lessonId":"a36a726f-6416-4956-a06e-4dc48e3f0eab"}



{
  "type": "exercise",
  "id": "fb122a1e-a8b0-4100-bc78-257da81ed885"
}

手算与调试时的检查清单

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

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

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

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

    Lz<0\frac{\partial L}{\partial z}<0

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

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

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


小结

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

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

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

  • 对线性节点:

z=wx+bz=wx+b

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

Lw=Lzx\frac{\partial L}{\partial w} = \frac{\partial L}{\partial z}x Lb=Lz\frac{\partial L}{\partial b} = \frac{\partial L}{\partial z}

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

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

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