Skip to content

反向传播与梯度优化 ​

标签
AI/ml/数学
字数
3440 字
阅读时间
15 分钟

前两层(03-线性代数与 GEMM、04-概率、Softmax 与信息论)描述的是前向:张量怎么变换、输出是什么分布。这一层描述参数怎么学 —— 链式法则沿计算图反向传播,以及训练过程中梯度为什么会消失、爆炸。

导数、梯度与 Jacobian ​

∇xf=[∂f∂x1⋮∂f∂xn],xt+1=xt−η∇f(xt)

梯度指向函数上升最快的方向,所以梯度下降沿负梯度走。

向量函数 y=f(x) 的 Jacobian 描述每个输入分量对每个输出分量的局部影响:

Jij=∂yi∂xj

标量函数的 Hessian 是二阶偏导矩阵,描述局部曲率。参数量太大时不会显式构造完整 Hessian,但 Hessian-vector product 与曲率近似仍用于优化与稳定性研究。

计算图与反向模式自动微分 ​

设 a=xy、b=a+x、L=b2。前向保存必要中间量,反向从 ∂L∂L=1 开始:

∂L∂b=2b,∂L∂a=2b

x 有两条路径影响 L,梯度必须相加:

∂L∂x=∂L∂b∂b∂x+∂L∂a∂a∂x=2b+2by,∂L∂y=2bx

这就是「一个张量被多处使用时梯度要累加」的原因 —— 不是实现细节,是链式法则里路径求和的直接结果。它也解释了 PyTorch 默认累加梯度的设计(见 01-PyTorch 框架与训练循环)。

概念上梯度是 Jacobian 连乘,但自动微分不会显式创建巨大 Jacobian,它算的是 vector-Jacobian product(反向模式)—— 一次反向就能拿到对全部参数的梯度,这是反向模式比前向模式适合神经网络的根本原因(参数多、输出少)。

一个具体计算图:a=xy、b=a+x、L=b2。

   前向(保存必要中间量)
     x ─┬──────────▶ b = a + x ──▶ L = b²
        │              ▲
        └──▶ a = x·y ──┘
        y ──▶

   反向(从 ∂L/∂L = 1 开始)
     ∂L/∂b = 2b
     ∂L/∂a = 2b
        │
        └─ x 有两条路径影响 L ⟹ 梯度必须相加
           ∂L/∂x = (∂L/∂b)(∂b/∂x) + (∂L/∂a)(∂a/∂x)
                 = 2b + 2b·y
           ∂L/∂y = 2b·x

   └─ 这就是「一个张量被多处使用时梯度要累加」的原因 ——
      不是实现细节,是链式法则里路径求和的直接结果。
      它也解释了 PyTorch 默认累加梯度的设计。

   概念上梯度是 Jacobian 连乘,但自动微分不会显式创建巨大 Jacobian
     它算的是 vector-Jacobian product(反向模式)
     └─ 一次反向就能拿到对全部参数的梯度 ——
        这是反向模式比前向模式适合神经网络的根本原因
        (参数多、输出少)

三个必须会手推的梯度 ​

这是全章最实用的部分 —— 面试与实际调试都会用到。

线性层 ​

Y=XW+b,X∈RB×din, W∈Rdin×dout

给定上游梯度 G=∂L∂Y∈RB×dout:

∂L∂X=GW⊤,∂L∂W=X⊤G,∂L∂b=∑i=1BGi,:

只看 shape 就能检查公式(这是最实用的自检手段):

三件与前向、反向配套的事,都有可算的判据

   ① 线性层反向的 shape 自检(只看 shape 就能检查公式)
        dX: (B, dout) @ (dout, din) → (B, din)
        dW: (din, B)  @ (B, dout)   → (din, dout)
        db: 对 batch 维归约          → (dout,)
        └─ 一个线性层的反向包含两次 GEMM 和一次归约 ——
           这是训练算力显著高于单次前向的原因:前向一次 GEMM,反向两次

   ② 初始化与方差传播
        y = Σ w_i x_i(w、x 独立且均值为 0)
        ⟹ Var(y) = n · Var(w) · Var(x)
        └─ 要让各层激活方差稳定,令 Var(w) 与 1/n 同阶 ——
           这就是 Xavier / Kaiming 初始化按 fan-in、fan-out
           与激活函数选尺度的来源
        └─ 初始化不是越小越稳定:
           过小让信号与梯度衰减,过大导致饱和或爆炸

   ③ 残差连接为什么帮助梯度传播
        y = x + F(x)  ⟹  ∂y/∂x = I + ∂F/∂x
        └─ 恒等路径提供了一条不完全依赖 F 的梯度通路
        └─ 它不保证永远不会消失 / 爆炸,
           但显著改善深层网络的优化条件

   症状对照:梯度范数趋近 0(早期层几乎不更新)/
             突然巨大(loss 变 inf / nan)/
             混合精度下更早触发下溢或溢出
   缓解手段:合理初始化、残差连接、归一化、梯度裁剪、
             合适的学习率、稳定的 dtype

一个线性层的反向包含两次 GEMM 和一次归约。 这是训练算力显著高于单次前向的原因 —— 前向一次 GEMM,反向两次。

Softmax 的 Jacobian ​

∂pi∂zj=pi(δij−pj),Jsoftmax=diag(p)−pp⊤

Softmax 每个输出依赖所有输入,所以 Jacobian 不是对角矩阵。但实际反向仍可在 O(V) 时间内算完,不需要物化 V×V 矩阵。

Softmax + 交叉熵 ​

one-hot 标签下,两者联合求导后得到深度学习里最重要的简化之一:

∂L∂zi=pi−yi

对正确类别梯度是 py−1,其他类别是 pi —— 预测越错,推动 logits 修正的力度越大。这个形式的好处在于它把「Softmax 的 Jacobian(V×V)」和「交叉熵的导数」两件麻烦事一起消掉了。

融合的交叉熵 kernel 可以共享最大值与归约结果、减少中间概率张量,并直接产出所需梯度 —— 这也是标准实现都提供 cross_entropy 而不让你手动拼 softmax + log 的原因。

softmax 与交叉熵联合求导后,得到深度学习里最重要的简化之一。

   单独看两者都很麻烦
     Softmax 的 Jacobian   J = diag(p) − p p^T   —— 每个输出依赖所有输入,
                                                 不是对角矩阵,是 V × V
     交叉熵的导数           要对 p 逐项求导

   联合起来:∂L/∂z_i = p_i − y_i

     ├─ 对正确类别    p_y − 1
     └─ 其他类别      p_i
     └─ 预测越错,推动 logits 修正的力度越大

   └─ 这个形式把「Softmax 的 Jacobian(V × V)」和「交叉熵的导数」
      两件麻烦事一起消掉了。

   工程上的推论
     融合的交叉熵 kernel 可以共享最大值与归约结果、
     减少中间概率张量,并直接产出所需梯度 ——
     这也是标准实现都提供 cross_entropy,
     而不让你手动拼 softmax + log 的原因。

   梯度检查(小规模自检)
     中心差分:(∂f/∂x_i) ≈ [f(x + εe_i) − f(x − εe_i)] / (2ε)
     ε 不能太大(截断误差)也不能太小(被浮点舍入淹没)
     └─ 适合小尺寸、双精度、确定性函数;
        不适合直接检查带随机性或巨大张量的完整训练任务

梯度检查 ​

自动微分实现可用中心差分做小规模检查:

∂f∂xi≈f(x+ϵei)−f(x−ϵei)2ϵ

ϵ 不能太大(截断误差)也不能太小(被浮点舍入淹没)。适合小尺寸、双精度、确定性函数,不适合直接检查带随机性或巨大张量的完整训练任务。

优化:从 SGD 到 Adam ​

g^=1B∑i∈B∇θLi(θ),θt+1=θt−ηg^t

mini-batch 梯度是全量梯度的随机估计。更大 batch 降低梯度估计方差并提高硬件利用率,但需要更多激活内存,且超参数可能要相应调整。

动量 对历史梯度做指数加权平均:

vt=βvt−1+(1−β)gt,θt+1=θt−ηvt

Adam 同时维护一阶矩与二阶矩,偏差修正后逐元素自适应缩放:

mt=β1mt−1+(1−β1)gt,vt=β2vt−1+(1−β2)gt2θt+1=θt−ηm^tv^t+ϵ

Adam 的代价是状态 —— 每个参数额外保存 m 与 v 两份。若状态用 FP32,训练显存的大头就是这三项(主权重 + m + v = 12Ψ)。这正是 ZeRO 第一刀切优化器状态的原因,账本见 02-分布式训练总论与显存账本。

Adam 的代价是状态。

   ┌─ SGD ────────────────────────────────────────────────┐
   │ θ ← θ − η · ĝ                                           │
   │ 状态:无(只需要梯度本身)                                 │
   └──────────────────────────────────────────────────────┘
   ┌─ 动量 ───────────────────────────────────────────────┐
   │ v_t = β v_{t−1} + (1 − β) g_t                           │
   │ θ ← θ − η v_t                                           │
   │ 状态:一份 v                                             │
   └──────────────────────────────────────────────────────┘
   ┌─ Adam ───────────────────────────────────────────────┐
   │ m_t = β₁ m_{t−1} + (1 − β₁) g_t          一阶矩         │
   │ v_t = β₂ v_{t−1} + (1 − β₂) g_t²         二阶矩         │
   │ 偏差修正后逐元素自适应缩放                                 │
   │ θ ← θ − η · m̂_t / (√v̂_t + ε)                            │
   │ 状态:两份(m 与 v)                                      │
   └──────────────────────────────────────────────────────┘

   显存账本:若状态用 FP32,每个参数要 4 + 4 + 4 = 12 字节
     主权重 + m + v = 12Ψ
     └─ 训练显存的大头就是这三项 ——
        这正是 ZeRO 第一刀切优化器状态的原因。

   mini-batch 梯度是全量梯度的随机估计
     更大 batch 降低梯度估计方差并提高硬件利用率,
     但需要更多激活内存,且超参数可能要相应调整。

梯度消失与爆炸 ​

深层复合函数的梯度是 Jacobian 连乘:

∂L∂h0=∂L∂hL∏l=1L∂hl∂hl−1

若这些变换在相关方向上的尺度长期小于 1,梯度指数衰减;长期大于 1,指数增长。

工程症状:

  • 梯度范数趋近 0,早期层几乎不更新
  • 梯度范数突然巨大,loss 变成 inf / nan
  • 混合精度下更早触发下溢或溢出

缓解手段:合理初始化、残差连接、归一化、梯度裁剪、合适的学习率、稳定的 dtype。

残差连接为什么帮助梯度传播 ​

y=x+F(x)⟹∂y∂x=I+∂F∂x

恒等路径提供了一条不完全依赖 F 的梯度通路。它不保证永远不会消失/爆炸,但显著改善深层网络的优化条件。

初始化与方差传播 ​

y=∑i=1nwixi⟹Var(y)=nVar(w)Var(x)

(wi,xi 独立且均值为 0。)要让各层激活方差稳定,令 Var(w) 与 1/n 同阶 —— 这就是 Xavier / Kaiming 初始化按 fan-in、fan-out 与激活函数选尺度的来源。

初始化不是越小越稳定:过小让信号与梯度衰减,过大导致饱和或爆炸。

归一化:LayerNorm 与 RMSNorm ​

对一个 token 的隐藏向量 x∈RH:

μ=1H∑ixi,σ2=1H∑i(xi−μ)2,LayerNorm(xi)=γixi−μσ2+ϵ+βi

LayerNorm 沿隐藏维归约,不依赖 batch 中其他样本 —— 这是它适合可变序列与自回归模型的原因(BatchNorm 做不到)。

从 kernel 角度,它至少涉及均值归约、方差归约、逐元素变换三步。朴素实现多次读写 HBM;优化实现会把归约与仿射变换融合,并处理长向量的并行归约。注意方差用的是总体形式(除以 H),见 04-概率、Softmax 与信息论。

RMSNorm 不减均值,只按均方根缩放:

RMS(x)=1H∑ixi2+ϵ,RMSNorm(xi)=γixiRMS(x)

减少了均值相关计算 —— 但**「公式更少」不必然按同比例提升端到端速度**,实际收益取决于融合程度、访存模式与模型整体瓶颈在哪。

更深入的结构与工程取舍见 13-归一化与残差连接。

梯度裁剪 ​

g←g⋅min(1,c‖g‖2+ϵ)

超过阈值 c 时整体等比缩小,方向不变。

梯度裁剪在分布式下会引入一次集合通信

若梯度被分片(ZeRO / FSDP),计算全局 L2 范数需要跨 rank 汇总局部平方和:

‖g‖2=∑r∑i∈rgi2

一个看起来纯本地的数学操作,在分片场景下变成了一次 AllReduce。这是「分布式下算不准」的常见来源 —— 只对本地分片裁剪,等价于用错误的范数做了裁剪。

裁剪本身只有一个式子,但梯度被分片后它就不再是就地操作了。

   梯度裁剪:g ← g · min(1, c / (‖g‖₂ + ε))
        └─ 超过阈值 c 时整体等比缩小,方向不变

   ┌─ 单卡 ──────────────────────────────────────────────┐
   │ 算全局 L2 范数 ──▶ 就地裁剪 ──▶ 继续                 │
   └────────────────────────────────────────────────────┘
   ┌─ 梯度被分片(ZeRO / FSDP)──────────────────────────┐
   │ 每个 rank 只有自己那一片的梯度                         │
   │ ‖g‖₂ = √( Σ_r Σ_{i ∈ r} g_i² )                       │
   │        └─ 局部平方和必须先跨 rank 汇总 ──▶ 一次 AllReduce │
   │ 拿到全局范数后才能各自裁剪                             │
   └────────────────────────────────────────────────────┘

   └─ 这是「分布式下算不准」的常见来源:
      只对本地分片裁剪,等价于用错误的范数做了裁剪。

相关 ​

参考 ​

贡献者 ​

文件历史 ​