机器学习数学基础 120 章

多元链式法则、计算图与反向传播

层级:A|必学

1. 反向传播不是神秘的专用公式

它把复杂函数拆成简单节点,记录前向中间值,再从标量损失反向应用链式法则。关键有两条:沿一条路径导数相乘;同一变量到输出有多条路径时贡献相加。

2. 多元标量链式法则

z=f(u1,ldots,um),uj=uj(x).z=f(u_1,ldots,u_m), \qquad u_j=u_j(x).

dzdx=j=1mzujdujdx.\frac{dz}{dx} =\sum_{j=1}^{m} \frac{\partial z}{\partial u_j} \frac{du_j}{dx}.

求和来自 xx 通过每个 uju_j 影响 zz 的所有路径。

例:z=u2+uvz=u^2+uvu=x2u=x^2v=3xv=3x

dzdx=(2u+v)(2x)+u(3).\frac{dz}{dx} =(2u+v)(2x)+u(3).

3. 向量形式

y=f(x)y=f(x),标量损失 L=L(y)L=L(y)

xL=Jf(x)TyL.\nabla_xL=J_f(x)^T\nabla_yL.

这表示把输出空间的梯度通过局部 Jacobian 的转置拉回输入空间。多层复合:

x0L=J1TJ2TJkTxkL,\nabla_{x_0}L =J_1^TJ_2^T\cdots J_k^T\nabla_{x_k}L,

实际逐层计算,避免显式乘出巨大 Jacobian。

4. 计算图例子

单样本线性单元:

z=wTx+b,p=σ(z),L=[ylogp+(1y)log(1p)].z=w^Tx+b, \qquad p=\sigma(z), \qquad L=-[y\log p+(1-y)\log(1-p)].

节点依赖为

(w,x,b)zpL.(w,x,b)\to z\to p\to L.

前向计算保存 z,pz,p;反向从 L/L=1\partial L/\partial L=1 开始。

5. Sigmoid + 二元交叉熵的简化

先求

Lp=yp+1y1p,\frac{\partial L}{\partial p} =-\frac yp+\frac{1-y}{1-p}, dpdz=p(1p).\frac{dp}{dz}=p(1-p).

相乘:

Lz=(yp+1y1p)p(1p)=py.\frac{\partial L}{\partial z} =\left(-\frac yp+\frac{1-y}{1-p}\right)p(1-p) =p-y.

再由 z=wTx+bz=w^Tx+b

wL=(py)x,\nabla_wL=(p-y)x, Lb=py,\frac{\partial L}{\partial b}=p-y, xL=(py)w.\nabla_xL=(p-y)w.

合并 logits 损失不仅公式简洁,而且能用稳定的 LogSumExp 技巧避免 pp 舍入为 0 或 1。

6. 分支节点的梯度相加

y=f(x)+g(x),y=f(x)+g(x),

dydx=f(x)+g(x).\frac{dy}{dx}=f'(x)+g'(x).

在图中 xx 分叉到两条路径,反向汇合时把梯度加总。残差连接

y=x+F(x)y=x+F(x)

的梯度为

Lx=Ly(I+JF(x)),\frac{\partial L}{\partial x} =\frac{\partial L}{\partial y} \left(I+J_F(x)\right),

恒等路径提供直接梯度通道。

7. 批量全连接层

前向:

Z=XW+1bT,Z=XW+\boldsymbol1b^T,

上游梯度 G=L/ZG=\partial L/\partial Z。反向:

LW=XTG,\frac{\partial L}{\partial W}=X^TG, LX=GWT,\frac{\partial L}{\partial X}=GW^T, Lb=GT1\frac{\partial L}{\partial b}=G^T\boldsymbol1

(等价于沿 batch 维求和)。每个公式都可通过形状检查:WW 的梯度必须与 WW 同形状。

8. 激活函数的逐元素反传

H=ϕ(Z)H=\phi(Z) 逐元素作用,上游为 GHG_H

GZ=GHϕ(Z).G_Z=G_H\odot\phi'(Z).

\odot 是 Hadamard 乘。ReLU 导数在正区间为 1、负区间为 0,0 点由框架约定。若某神经元长期落在负区间,梯度为 0,可能出现“死亡 ReLU”。

9. 反向模式为何高效

若函数有 pp 个输入参数、一个标量输出:

  • 逐参数数值差分需要约 pp 次前向;
  • 前向模式通常每个方向传播一次;
  • 反向模式一次前向加一次反向,就得到所有 pp 个偏导。

代价通常是常数倍前向计算,但需要保存或重算中间值,形成时间—内存权衡。梯度检查只在小规模参数或随机方向上做。

10. 自动微分的三种误解

  1. 不是符号化简:它沿具体计算图组合局部导数;
  2. 不是数值差分:除浮点舍入外通常给出精确链式导数;
  3. 不是自动理解数学意图:detach、原地修改、错误广播和不稳定公式仍会产生错误或无用梯度。

11. 梯度消失与爆炸

长链式乘积中,Jacobians 的奇异值控制梯度尺度。许多小于 1 的因子导致消失,许多大于 1 导致爆炸。常见缓解:

  • 合理初始化;
  • ReLU/GELU 等激活;
  • 残差连接;
  • 归一化;
  • 门控结构;
  • 梯度裁剪。

12. 梯度累积与清零

许多框架默认把新梯度累加到参数 .grad,对应多个损失对同一参数的贡献相加。每个优化步骤前若不清零,会无意累加跨批梯度。主动梯度累积则用多个微批近似大 batch,需确认损失是求和还是平均。

易错点

  1. 分支路径在反向汇合时要相加。
  2. 逐元素乘与矩阵乘必须区分。
  3. 上游梯度的形状决定局部 VJP,不应显式构造 Jacobian。
  4. detach 会切断计算图。
  5. 数值稳定的合并损失优于先算概率再取对数。

常见问答

Q1:反向传播会不会更新参数?

反向传播只计算梯度;优化器根据梯度、学习率和状态更新参数。框架 API 有时把两步紧邻调用,但概念不同。

Q2:为什么 bias 梯度要沿 batch 求和?

同一个偏置被广播到每个样本,每条使用路径都对它产生贡献,链式法则要求相加。

Q3:训练时能否不保存激活以省内存?

可用 gradient checkpointing,只保存部分节点,反向时重算其余前向,牺牲计算换内存。

练习

  1. y=x2+x3y=x^2+x^3 用路径求和解释 dy/dxdy/dx
  2. z=wx+b,L=z2z=wx+b,L=z^2w,b,xw,b,x 梯度。
  3. 写出 y=x+F(x)y=x+F(x) 的导数。
  4. XXn×dn\times dWWd×hd\times hGGn×hn\times h,检查 XTGX^TGGWTGW^T 形状。
  5. 为什么 BCEWithLogitssigmoid 后再 log 稳定?

答案与提示

  1. 两路径贡献 2x2x3x23x^2 相加。
  2. 2zx,2z,2zw2zx,2z,2zw
  3. I+JFI+J_F,标量时为 1+F(x)1+F'(x)
  4. 分别为 d×hd\times hn×dn\times d
  5. 可代数合并并用 LogSumExp/softplus 形式,避免指数溢出和 log0\log0