40 分钟
AI 数学精要

链式法则与计算图:反向传播的本质

亲手对复合函数和最小计算图做前向与反向,彻底搞懂反向传播就是沿图反向套用链式法则

  • 会用链式法则求复合函数的导数,还能手算到具体数值
  • 能画出最小计算图,区分前向算值与反向算梯度
  • 理解路径相乘、分叉相加的多元链式规则
  • 讲清楚:参数很多、但损失值只有一个的时候,反向模式为什么比前向模式效率高

反向传播没有魔法,它就是「链式法则 + 一张计算图」

神经网络本质是个层层嵌套的巨型复合函数:输入过线性层、激活、注意力……最后输出一个损失标量。要训练它,就得求损失对每个参数的导数。直接对这个嵌套的复杂结构求导根本没法下手,链式法则能把它拆成每个简单运算的局部导数,再顺着数据流反向串起来就行。这一节你得亲手走通「前向算数值、反向传梯度」的全过程——走完你就明白,PyTorch 的 autograd 干的就是把这件事自动化了。

2.1 复合函数与链式法则:层层相乘

复合函数求导的链式法则,直接说公式:如果 y=f(u)、u=g(x),也就是 y=f(g(x)),那么 dy/dx = (dy/du)·(du/dx)。就是外层函数对中间变量的导数,乘中间变量对输入的导数。要是多层嵌套,就一路乘到底。 直觉上就是「传导」:x 变一点先带动 u 变,u 的变化再经过放大或者缩小传给 y,总的变化效应就是每一段变化率的乘积。

手算 y=(x²+1)³。令 u=x²+1,则 y=u³。外层 dy/du=3u²,内层 du/dx=2x,所以 dy/dx=3u²·2x=6x(x²+1)²。代入 x=2:u=5,dy/dx=6×2×5²=12×25=300。你可以用数值差分核验:x=2.001 时 y=(5.004001)³≈125.300,与 x=2 时的 125 相差约 0.300,除以 0.001 得约 300,一致。

示例代码(可运行)
填空题填写空白处的代码
y = (x²+1)³,令 u=x²+1 # dy/du = # 填 3*u**2 # du/dx = # 填 2*x # dy/dx = 3u²·2x,x=2 时 = # 填 300

链可以更长。对三重复合 y=f(g(h(x))),dy/dx=f′(g(h))·g′(h)·h′(x),逐环相乘。遇到看似复杂的嵌套,只需从外到内逐层问「这一层对它的直接输入导数是多少」再连乘。反向传播正是把这套动作机械化:不管网络多少层,每个节点只负责自己这一环的局部导数,全局导数由计算图自动串出,人不需要对整个网络写一个巨型导数。

2.2 多元链式:路径相乘,分叉相加

中间变量有多个的时候,规则是「沿每条路径相乘,跨多条路径相加」。比如 z=f(u,v),u、v 都依赖 x,那 dz/dx = (∂z/∂u)·(du/dx) + (∂z/∂v)·(dv/dx)。 为什么要相加?x 对 z 的影响走两条独立的通道,总效应就是两条通道的和加起来。 这条「分叉要相加」是反向传播里最容易漏的环节——同一个中间量被后面好几个地方用时,它会收到各条支路传回来的梯度,必须累加。

选择题

中间变量 u 被后面两个不同分支同时用了,反向传播的时候它会收到两份梯度,正确处理方式是什么?

2.3 计算图:节点是运算,边是数据流

计算图把复合函数显式画成有向无环图:节点是变量或一次运算结果,边表示「谁由谁算出」。前向传播从输入出发,沿箭头算出每个节点的值并缓存。反向传播从损失标量出发,逆着箭头,用链式法则逐节点算梯度。每个节点只需知道两件事:前向怎么由输入算输出,以及局部导数(输出对各输入的导数),全局梯度由框架自动拼接。

走一个最小计算图(线性单元加平方损失):输入 x=2、参数 w=3、b=1,目标 t=1。前向:z=w·x+b=3×2+1=7;损失 L=0.5·(z−t)²=0.5×6²=18。反向(从 L 往回):dL/dz = z−t = 6(因为 0.5(z−t)² 导数就是 z−t)z=wx+b 的局部导数:∂z/∂w=x=2、∂z/∂x=w=3、∂z/∂b=1dL/dw = (dL/dz)·(∂z/∂w) = 6×2 = 12dL/dx = 6×3 = 18,dL/db = 6×1 = 6于是 w 的梯度是 12、b 的梯度是 6,更新时 w←w−η×12、b←b−η×6,损失就会下降。整个神经网络只是把这张图放大到亿万节点。

示例代码(可运行)

再走一个带激活的两层最小例。标量输入 x=1,第一层 w1=2、b1=0:z1=w1·x=2,激活 a=σ(z1)=1/(1+e⁻²)≈0.881;第二层 w2=−1:z2=w2·a≈−0.881,损失 L=0.5·z2²≈0.388。反向:dL/dz2=z2≈−0.881;dL/dw2=(dL/dz2)·a≈−0.776;dL/da=(dL/dz2)·w2≈0.881;而 σ′(z1)=a(1−a)≈0.881×0.119≈0.105,故 dL/dz1=(dL/da)·σ′≈0.0925,dL/dw1=(dL/dz1)·x≈0.0925。注意梯度穿过一次 sigmoid 就从 0.881 缩到 0.0925、缩小近十倍——多层 sigmoid 连乘正是梯度消失的第一现场。

预测输出
同一计算图,若把 w 改成 4(x=2,b=1,t=1),则 z=9、dL/dz=8,此时 dL/dw 等于多少?

2.4 前向模式 vs 反向模式:为什么深度学习选反向

自动微分有两种扫图方向。前向模式:从输入开始,随前向同时传递每个变量对「某一个输入」的导数,要对 n 个输入求导就得扫 n 遍。反向模式:先完整前向一遍,再从输出反向扫一遍,一次就能同时得到这一个输出对「全部输入」的导数。

深度学习的结构刚好是「输入/参数极多(上亿)、输出极少(一个损失标量)」。用前向模式得为上亿参数各扫一遍,代价完全扛不住;反向模式只扫两遍(一遍前向、一遍反向)就能拿到所有参数梯度,复杂度和一次前向是同阶的。这就是反向传播能当训练基石的根本原因——不是它更「精确」,而是它在这种输入多输出少的结构下,计算量能差好几个数量级,省太多了。

前向传播的时候,同时把对单个输入的导数传过去。输入有 n 个就得扫 n 遍,适合输出多、输入少的场景,比如实验设计、部分物理仿真。

给你算笔账,看反向模式到底要花多少代价:假设一次前向的计算量是1,反向得给每个前向算子再算局部导数、做乘法,前向加反向的总计算量大概是一次前向的2倍——这就是「训练一步约等于两次推理」的由来。代价是得占显存存前向的激活值,留着反向的时候复用,说白了就是用存储换计算。stage-14 里的激活重计算(也就是梯度检查点)走另一个路子:反向的时候把前向再跑一遍,用多出来的约三分之一计算量,换显存大幅下降。拿全连接层 X[m,k]@W[k,n] 举例子,前向大概要 2mnk 次乘加,反向得分别对 W 和 X 求梯度,每样又各约 2mnk,加起来刚好是前向的两倍,和上面的总量估算对得上。

推导

经典局部导数:sigmoid 为什么能写成自指形式

σ(x)=1/(1+e⁻ˣ),直接求导挺麻烦,但结论用得特别多:σ′(x)=σ(x)·(1−σ(x))。推导过程:σ′=e⁻ˣ/(1+e⁻ˣ)²=σ·[e⁻ˣ/(1+e⁻ˣ)]=σ·(1−σ)。 这意味着反向传播时,只要把前向传播算出来的 σ 缓存下来,就不用再算指数了,直接用 σ(1−σ) 就能得到导数。 手算验证下:x=0 时 σ(0)=0.5,σ′(0)=0.5×0.5=0.25;x=2 时 σ≈0.881,σ′≈0.881×0.119≈0.105。能看出来,|x| 一大,导数就趋近 0——这正是 sigmoid 深层梯度消失的直接来源,链式法则连乘一串接近 0 的数,梯度会迅速衰减。

2.5 计算图工程:缓存、累加与切断

反向传播能不能跑对、跑快,就看三个工程细节。前向传播时得把中间激活值缓存下来,留着反向的时候用——这也是训练比纯推理费显存的原因(stage-14 讲的激活重计算,就是用算力换这块显存开销)。同一个计算节点被多次复用的话,梯度必须累加,不能直接覆盖。可以主动切断计算图:让某段计算不参与反向传播(推理场景用 no_grad、要停掉某个张量的梯度用 detach),用来冻结参数或者阻断没必要的梯度回传。

找 Bug被多个分支用到的节点梯度,得按链式法则加起来,直接覆盖就会丢一半梯度;不调 zero_grad 的话,上一个 batch 的梯度会累进来,相当于偷偷把 batch 搞大了。
# 一个节点 h 被两个输出分支使用,反传时写成「覆盖」 grad_h = grad_branch1 # 只留了第一路 # grad_h = grad_branch2 # 第二路被注释丢弃 # 训练循环里多个 batch 反传前也没清空旧梯度 loss.backward() optimizer.step() # 没有 optimizer.zero_grad()
配对题把计算图环节对到它的作用

实操里有三个高频的计算图踩坑点,挨个说:in-place 原地改写(比如 x+=1、张量下标赋值)会改掉反向传播要用到的前向值,框架要么报错要么给错误梯度。写成 x=x+1 生成新张量就行。中途把张量转成 NumPy 再转回来,会直接切断计算图——NumPy 本身不带自动微分信息。跨框架操作要用对应框架的算子。默认只有叶子张量(用户自己创建、且 requires_grad 为 True 的参数)会保留 .grad 属性,中间张量的梯度用完就被释放了。要观察中间张量的梯度,得显式调用 retain_grad。另外每轮反向传播前要 zero_grad,不然梯度会跨 batch 累加。

ℹ️雅可比:向量值函数的偏导表

输出是向量而非标量时(一层输出多个分量),它对输入的全部偏导排成雅可比矩阵 J,Jᵢⱼ=∂yᵢ/∂xⱼ。反向传播本质是让上游梯度逐个乘上每个算子的雅可比。逐元素运算(ReLU、平方)的雅可比是对角阵、计算极简,这也是它们高效的原因;矩阵乘的雅可比会还原成另一侧矩阵,解释了 dL/dW=(dL/dZ)·Xᵀ 的形状来源。

🐍autograd 的本质与常见面试点

PyTorch 的 autograd 就是动态构建计算图,backward() 时做反向模式自动微分:张量设 requires_grad 后,会记录自己由哪些运算生成(就是 grad_fn 链),反向时按拓扑逆序调用每个算子的 backward。面试常问的几个点:训练时显存为啥随 batch 或序列长度涨?因为缓存了激活值。retain_graph 什么时候用?要对同一个计算图做多次反向的时候。view 之后反向为啥没问题?框架记了形状变换的映射关系。

⚠️梯度消失/爆炸的链式乘积根源

链式法则本质就是连乘。每层局部导数要是一直小于 1——比如 sigmoid 导数≤0.25——多层乘下来会指数级趋近于 0,这就是梯度消失,浅层网络直接学不动。反过来,导数一直大于 1 就会指数级放大,叫梯度爆炸,权重直接变 NaN。残差连接、合理初始化、LayerNorm、梯度裁剪,分别从四个方向解决:加直通路径、稳住单层方差、归一化、截断。

💡手推梯度的固定套路

遇到陌生层别硬背:写出前向标量表达式;从损失写 dL/dout;对每个输入写局部导数;相乘得分支、分叉就相加;用中心差分对一个样例数值核对。五步走完,再复杂的注意力层也能推。

选择题

深度学习用反向模式自动微分,最主要的原因是?

本节小结

一条因果链:复合函数靠链式法则「外层导数×内层导数」层层传导 → 多元情形路径相乘、分叉相加 → 计算图前向算值并缓存、反向从损失逆箭头套链式得到每个参数梯度 → 输入海量而损失唯一时反向模式两遍扫图远胜前向模式 → 工程上注意缓存激活、梯度累加、迭代清零与必要时切断图。反向传播=链式法则在计算图上的系统化执行。

资深工程师加餐

底层原理 · 大厂视角 · 工程经验,点卡片展开

信息熵度量一个分布的不确定程度,越随机熵越大。交叉熵衡量「用模型预测的分布去编码真实分布」所需的平均代价:模型预测越接近真实标签,交叉熵越小。因此多分类任务用交叉熵做损失不是拍脑袋,而是有严格信息论依据的最大似然等价形式。