反向传播与梯度优化
前两层(03-线性代数与 GEMM、04-概率、Softmax 与信息论)描述的是前向:张量怎么变换、输出是什么分布。这一层描述参数怎么学 —— 链式法则沿计算图反向传播,以及训练过程中梯度为什么会消失、爆炸。
导数、梯度与 Jacobian
梯度指向函数上升最快的方向,所以梯度下降沿负梯度走。
向量函数
标量函数的 Hessian 是二阶偏导矩阵,描述局部曲率。参数量太大时不会显式构造完整 Hessian,但 Hessian-vector product 与曲率近似仍用于优化与稳定性研究。
计算图与反向模式自动微分
设
这就是「一个张量被多处使用时梯度要累加」的原因 —— 不是实现细节,是链式法则里路径求和的直接结果。它也解释了 PyTorch 默认累加梯度的设计(见 01-PyTorch 框架与训练循环)。
概念上梯度是 Jacobian 连乘,但自动微分不会显式创建巨大 Jacobian,它算的是 vector-Jacobian product(反向模式)—— 一次反向就能拿到对全部参数的梯度,这是反向模式比前向模式适合神经网络的根本原因(参数多、输出少)。
一个具体计算图:
前向(保存必要中间量)
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(反向模式)
└─ 一次反向就能拿到对全部参数的梯度 ——
这是反向模式比前向模式适合神经网络的根本原因
(参数多、输出少)三个必须会手推的梯度
这是全章最实用的部分 —— 面试与实际调试都会用到。
线性层
给定上游梯度
只看 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
Softmax 每个输出依赖所有输入,所以 Jacobian 不是对角矩阵。但实际反向仍可在
Softmax + 交叉熵
one-hot 标签下,两者联合求导后得到深度学习里最重要的简化之一:
对正确类别梯度是
融合的交叉熵 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ε)
ε 不能太大(截断误差)也不能太小(被浮点舍入淹没)
└─ 适合小尺寸、双精度、确定性函数;
不适合直接检查带随机性或巨大张量的完整训练任务梯度检查
自动微分实现可用中心差分做小规模检查:
优化:从 SGD 到 Adam
mini-batch 梯度是全量梯度的随机估计。更大 batch 降低梯度估计方差并提高硬件利用率,但需要更多激活内存,且超参数可能要相应调整。
动量 对历史梯度做指数加权平均:
Adam 同时维护一阶矩与二阶矩,偏差修正后逐元素自适应缩放:
Adam 的代价是状态 —— 每个参数额外保存
与 两份。若状态用 FP32,训练显存的大头就是这三项(主权重 + + = )。这正是 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 连乘:
若这些变换在相关方向上的尺度长期小于 1,梯度指数衰减;长期大于 1,指数增长。
工程症状:
- 梯度范数趋近 0,早期层几乎不更新
- 梯度范数突然巨大,loss 变成
inf/nan - 混合精度下更早触发下溢或溢出
缓解手段:合理初始化、残差连接、归一化、梯度裁剪、合适的学习率、稳定的 dtype。
残差连接为什么帮助梯度传播
恒等路径提供了一条不完全依赖
初始化与方差传播
(
初始化不是越小越稳定:过小让信号与梯度衰减,过大导致饱和或爆炸。
归一化:LayerNorm 与 RMSNorm
对一个 token 的隐藏向量
LayerNorm 沿隐藏维归约,不依赖 batch 中其他样本 —— 这是它适合可变序列与自回归模型的原因(BatchNorm 做不到)。
从 kernel 角度,它至少涉及均值归约、方差归约、逐元素变换三步。朴素实现多次读写 HBM;优化实现会把归约与仿射变换融合,并处理长向量的并行归约。注意方差用的是总体形式(除以
RMSNorm 不减均值,只按均方根缩放:
减少了均值相关计算 —— 但**「公式更少」不必然按同比例提升端到端速度**,实际收益取决于融合程度、访存模式与模型整体瓶颈在哪。
更深入的结构与工程取舍见 13-归一化与残差连接。
梯度裁剪
超过阈值
梯度裁剪在分布式下会引入一次集合通信
若梯度被分片(ZeRO / FSDP),计算全局 L2 范数需要跨 rank 汇总局部平方和:
一个看起来纯本地的数学操作,在分片场景下变成了一次 AllReduce。这是「分布式下算不准」的常见来源 —— 只对本地分片裁剪,等价于用错误的范数做了裁剪。
裁剪本身只有一个式子,但梯度被分片后它就不再是就地操作了。
梯度裁剪:g ← g · min(1, c / (‖g‖₂ + ε))
└─ 超过阈值 c 时整体等比缩小,方向不变
┌─ 单卡 ──────────────────────────────────────────────┐
│ 算全局 L2 范数 ──▶ 就地裁剪 ──▶ 继续 │
└────────────────────────────────────────────────────┘
┌─ 梯度被分片(ZeRO / FSDP)──────────────────────────┐
│ 每个 rank 只有自己那一片的梯度 │
│ ‖g‖₂ = √( Σ_r Σ_{i ∈ r} g_i² ) │
│ └─ 局部平方和必须先跨 rank 汇总 ──▶ 一次 AllReduce │
│ 拿到全局范数后才能各自裁剪 │
└────────────────────────────────────────────────────┘
└─ 这是「分布式下算不准」的常见来源:
只对本地分片裁剪,等价于用错误的范数做了裁剪。相关
- 01-机器学习介绍 —— 参数与超参数的分界(谁在训练中被更新)
- 02-线性回归 —— 学习率上界与条件数的完整推导
- 03-线性代数与 GEMM —— 线性层反向为什么还是矩阵乘
- 04-概率、Softmax 与信息论 —— 交叉熵与 KL 的定义
- 01-PyTorch 框架与训练循环 —— autograd 与梯度累加的工程表现
- 01-数值计算与精度 —— Loss Scaling、溢出下溢与 FP32 主权重
- 02-分布式训练总论与显存账本 —— Adam 状态占了训练显存的多少
- 11-Self-Attention 机制 —— 除以
的方差推导
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块一-前置知识/第2章-数学基础
- https://jmlr.org/papers/v18/17-468.html
- https://arxiv.org/abs/1802.01528
- https://arxiv.org/abs/1412.6980
- https://arxiv.org/abs/1607.06450
- https://arxiv.org/abs/1910.07467
YJ