反向传播解决了“怎样计算梯度”,但训练还没有完成。拿到梯度以后,仍然要回答几个工程上极其实际的问题:一次梯度该看多少数据?不同参数的梯度量级相差很大怎么办?网络堆深以后,梯度为什么会消失或爆炸?隐藏状态又该怎样保持稳定?
这一章将这些问题串成一条训练链路:
1
2
3
4
5
6
7
mini-batch 给出带噪声的梯度估计
↓
SGD / Momentum / Adam 决定怎样更新参数
↓
残差、初始化、归一化让深层信号和梯度保持可用
↓
梯度裁剪处理偶发的极端更新
1. mini-batch:一次更新该使用多少数据?
设模型的全部可训练参数为 (\theta),训练集有 (N) 条样本。第 (i) 条样本的损失为:
$L_i(\theta)$
整个训练集上的平均损失为:
$L(\theta)=\dfrac{1}{N}\sum_{i=1}^{N}L_i(\theta)$
若每次更新都使用全部样本,梯度是:
$\nabla_\theta L(\theta)=\dfrac{1}{N}\sum_{i=1}^{N}\nabla_\theta L_i(\theta)$
这叫全量梯度。它很准确,但若数据集有数千万条样本,为更新一次参数而完整跑一遍数据集,成本太高。
实际训练每次随机取出一个 mini-batch (\mathcal B),其样本数为 (B):
$B=|\mathcal B|$
并用 batch 内梯度平均值估计全量梯度:
$\hat g=\dfrac{1}{B}\sum_{i\in\mathcal B}\nabla_\theta L_i(\theta)$
最基本的 SGD 更新是:
$\theta_{t+1}=\theta_t-\eta\hat g_t$
其中 (\eta) 是学习率。
一个数字例子
假设当前某个参数处,6 条训练样本分别给出下面的梯度:
1
2
样本编号 1 2 3 4 5 6
单样本梯度 2 4 6 8 10 12
全量梯度为:
$\dfrac{2+4+6+8+10+12}{6}=7$
若当前 mini-batch 抽到样本 3、4,则:
$\hat g=\dfrac{6+8}{2}=7$
恰好与全量梯度一致。若抽到样本 1、2:
$\hat g=\dfrac{2+4}{2}=3$
它不精确,却仍然给出“参数应减小”的方向。每一步只看一小批数据,因此梯度会带噪声;很多步结合起来,整体方向会接近全量梯度。
1
2
3
小 batch:显存较低,梯度噪声较大
大 batch:梯度更平滑,激活显存和单步成本更高
全量数据:准确,但大模型训练中通常不可行
语言模型中,batch 往往包含 (B) 条序列,每条有 (S) 个 token。一次反向传播参与 loss 平均的数据量大致与 (B\times S) 有关。
2. Momentum:不要被一个 batch 的噪声带偏
普通 SGD 只看当前梯度:
$\theta_{t+1}=\theta_t-\eta g_t$
若不同 mini-batch 给出略有冲突的方向,参数会左右摇摆。动量通过历史梯度的指数移动平均来平滑这个过程:
$v_t=\beta v_{t-1}+(1-\beta)g_t$
$\theta_{t+1}=\theta_t-\eta v_t$
其中 (\beta) 通常取 0.9。
设连续两个 batch 的二维梯度为:
$g_1=\begin{bmatrix}1\\10\end{bmatrix},\qquad g_2=\begin{bmatrix}-1\\10\end{bmatrix}$
第一维正负交替,代表横向摇摆;第二维持续为正,代表整体一直应往同一方向走。令 (\beta=0.8)、(v_0=0),则:
$v_1=0.8v_0+0.2g_1=\begin{bmatrix}0.2\\2\end{bmatrix}$
$v_2=0.8v_1+0.2g_2=\begin{bmatrix}-0.04\\3.6\end{bmatrix}$
第一维的正负影响被抵消,第二维的持续方向被积累。因此动量的作用不是改变正确方向,而是:
$\boxed{\text{压低短期噪声,放大长期一致的方向}}$
公式中的 ((1-\beta)) 让 (v_t) 保持为加权平均。若梯度长期恒为 (g),那么 (v_t) 最终会趋近 (g),而不会无故放大梯度尺度。
3. Adam:每个参数都有自己的有效步长
动量解决了时间上的梯度噪声,但不同参数的梯度量级也可能相差极大。假设两个参数当前的梯度为:
$g=\begin{bmatrix}100\\0.01\end{bmatrix}$
若用学习率 (\eta=0.001) 的 SGD:
$\Delta\theta=-\eta g=\begin{bmatrix}-0.1\\-0.00001\end{bmatrix}$
第一个参数变化很大,第二个几乎不动。Adam 为每个参数维护两份状态。
第一份是一阶矩,记录平滑后的梯度方向:
$m_t=\beta_1m_{t-1}+(1-\beta_1)g_t$
第二份是二阶矩,记录梯度平方的典型大小:
$v_t=\beta_2v_{t-1}+(1-\beta_2)g_t^2$
平方是逐元素进行的。由于初始 (m_0=v_0=0) 会让前几步估计偏小,Adam 会做偏差修正:
$\hat m_t=\dfrac{m_t}{1-\beta_1^t},\qquad\hat v_t=\dfrac{v_t}{1-\beta_2^t}$
最终更新:
$\theta_{t+1}=\theta_t-\eta\dfrac{\hat m_t}{\sqrt{\hat v_t}+\epsilon}$
第一次更新时,忽略 (\epsilon) 并完成偏差修正后,近似有:
$\dfrac{\hat m_1}{\sqrt{\hat v_1}}\approx\begin{bmatrix}100/100\\0.01/0.01\end{bmatrix}=\begin{bmatrix}1\\1\end{bmatrix}$
这解释了 Adam 的直觉:梯度长期较大的参数会被其自身的二阶矩缩小步长,梯度较小的参数不会永远陷入几乎不更新的状态。
从 AI Infra 角度,Adam 也有显存代价。除参数和梯度外,它还要为每个参数保存:
1
2
m:一阶矩状态
v:二阶矩状态
大模型中,这些优化器状态常以 FP32 保存,因而成为 ZeRO 等状态分片方案的重要优化对象。
4. 梯度消失与爆炸:深层网络的连乘问题
对深度网络:
$h_0=x,\qquad h_l=f_l(h_{l-1})$
链式法则表明,早期层的梯度包含许多局部导数连乘:
$\dfrac{\partial L}{\partial h_0}=\dfrac{\partial L}{\partial h_L}\prod_{l=1}^{L}\dfrac{\partial h_l}{\partial h_{l-1}}$
若每层局部导数都约为 0.9,跨越 100 层后:
$0.9^{100}\approx0.000027$
梯度几乎消失,底层参数难以学习。若每层局部导数约为 1.1:
$1.1^{100}\approx13780$
梯度会快速放大,可能让 loss 变成 inf 或 nan。
向量层的严格写法中,这串局部导数就是 Jacobian 的连乘:
$\dfrac{\partial L}{\partial h_0}=J_1^TJ_2^T\cdots J_L^T\dfrac{\partial L}{\partial h_L}$
反向传播不显式构造这些 Jacobian,但逐层 VJP 连起来的数学效果等价于这串乘法。实际网络中,某些方向可能被压缩,另一些方向可能被放大,因此问题不是简单的“所有梯度一起变小或变大”。
5. 残差连接:给梯度保留一条直接通路
普通层写作:
$y=F(x)$
残差块则写作:
$\boxed{y=x+F(x)}$
其局部导数为:
$\dfrac{\partial y}{\partial x}=I+\dfrac{\partial F}{\partial x}$
其中 (I) 是单位矩阵,来自输入 (x) 直接加到输出上的恒等路径。
若某层 (F) 在当前位置的导数约为 0.1:
1
2
没有残差:dy/dx = 0.1
有残差: dy/dx = 1 + 0.1 = 1.1
因此梯度不必完全依赖复杂分支 (F)。从前向看,残差也把任务从“学习完整映射”改成“学习相对输入的修正”:当某层暂时无须改动信息时,只要学到 (F(x)\approx0),输出便近似为 (x)。
残差不是消除梯度问题的数学保证;若 (\partial F/\partial x) 恰好与恒等项抵消,仍然可能不稳定。但它显著改善了深层网络的优化条件,也是 Transformer 能堆叠大量模块的重要原因。
6. 初始化:训练开始时就要控制方差传播
考虑线性变换:
$y=\sum_{i=1}^{n}w_ix_i$
若输入和权重独立、均值都接近 0:
$\operatorname{Var}(y)=n\operatorname{Var}(w)\operatorname{Var}(x)$
这里 (n) 是输入维度,也叫 fan-in。
假设 (n=1000),且输入方差为 1。若权重标准差为 0.01,则:
$\operatorname{Var}(w)=0.01^2=0.0001$
$\operatorname{Var}(y)=1000\times0.0001=0.1$
信号逐层缩小。若权重标准差为 0.1:
$\operatorname{Var}(y)=1000\times0.1^2=10$
信号会逐层放大。
希望输出方差和输入方差大致相当,就需要:
$n\operatorname{Var}(w)\approx1$
$\boxed{\operatorname{Var}(w)\approx\dfrac1n,\qquad\operatorname{std}(w)\approx\dfrac1{\sqrt n}}$
这也是 Xavier / Glorot、Kaiming / He 初始化都将权重尺度与 fan-in、fan-out 联系起来的原因。前者常考虑输入、输出两端的方差,后者针对 ReLU 一类激活会采用不同尺度。
初始化还必须打破对称性。若两个神经元的权重全部相同,它们前向输出和反向梯度都会相同,之后也会永远相同。随机初始化既控制数值尺度,也让不同神经元学到不同特征。
7. LayerNorm:稳定每个 token 的隐藏尺度
设一个 token 的隐藏向量为:
$x=\begin{bmatrix}1\\2\\3\\4\end{bmatrix},\qquad H=4$
LayerNorm 沿隐藏维计算均值:
$\mu=\dfrac{1+2+3+4}{4}=2.5$
方差为:
$\sigma^2=\dfrac{(1-2.5)^2+(2-2.5)^2+(3-2.5)^2+(4-2.5)^2}{4}=1.25$
再标准化:
$\hat x_i=\dfrac{x_i-\mu}{\sqrt{\sigma^2+\epsilon}}$
忽略很小的 (\epsilon),得到:
$\hat x\approx\begin{bmatrix}-1.342\\-0.447\\0.447\\1.342\end{bmatrix}$
此时向量均值约为 0、方差约为 1。最终 LayerNorm 还会加上逐维可学习的缩放和平移:
$\operatorname{LayerNorm}(x_i)=\gamma_i\dfrac{x_i-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta_i$
对 Transformer hidden states:
$X\in\mathbb R^{B\times S\times H}$
LayerNorm 对每个 ((b,s)) 位置各自沿最后一维 (H) 归约:
1
2
3
4
输入 X : [B, S, H]
均值、方差 : [B, S]
LayerNorm 输出 : [B, S, H]
可学习 gamma,beta: [H]
它不依赖 batch 内其他样本,因此适合 batch 可变、序列长度可变、推理 batch 可能为 1 的语言模型。
8. RMSNorm:保留均值,只规范整体幅度
RMSNorm 不减均值,只计算均方根:
$\operatorname{RMS}(x)=\sqrt{\dfrac1H\sum_i x_i^2+\epsilon}$
$\operatorname{RMSNorm}(x_i)=\gamma_i\dfrac{x_i}{\operatorname{RMS}(x)}$
取:
$x=\begin{bmatrix}3\\4\end{bmatrix}$
则:
$\operatorname{RMS}(x)=\sqrt{\dfrac{3^2+4^2}{2}}=\sqrt{12.5}\approx3.536$
若 (\gamma=[1,1]),输出约为:
$\begin{bmatrix}0.849\\1.131\end{bmatrix}$
它的均值约为 0.99,而不是 0。也就是说,RMSNorm 限制的是向量整体大小,不强制去除整体偏移。
| 操作 | 减均值 | 缩放方式 | 常见可学习参数 |
|---|---|---|---|
| LayerNorm | 是 | 标准差 | (\gamma,\beta) |
| RMSNorm | 否 | 均方根 | 通常是 (\gamma) |
RMSNorm 的公式更简单,但端到端性能收益仍取决于归约、Kernel 融合、访存和模型整体瓶颈,而不只取决于少算了一次均值。
9. 梯度裁剪:把异常的大更新限制在阈值内
若某次反向传播产生异常大的梯度,优化器的一步更新可能直接破坏参数。全局 L2 范数裁剪先计算:
$\lVert g\rVert_2=\sqrt{\sum_i g_i^2}$
给定阈值 (c),再执行:
$g\leftarrow g\cdot\min\left(1,\dfrac{c}{\lVert g\rVert_2+\epsilon}\right)$
例如:
$g=\begin{bmatrix}6\\8\end{bmatrix},\qquad\lVert g\rVert_2=10$
若阈值 (c=5),缩放系数为 (5/10=0.5),所以:
$g_{clip}=0.5\begin{bmatrix}6\\8\end{bmatrix}=\begin{bmatrix}3\\4\end{bmatrix}$
新梯度的范数正好为 5,方向保持不变。若原梯度范数小于阈值,则缩放系数为 1,梯度完全不变。
1
2
3
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
裁剪应发生在 optimizer.step() 之前。它适合作为偶发梯度尖峰的安全阀,但不是根治方案:若梯度长期巨大,还需要检查学习率、初始化、数据、loss 和模型结构;若梯度消失,裁剪不会有帮助。
分布式训练中,模型梯度可能分散在多个 rank 上。全局范数需要汇总所有 rank 的局部平方和:
$\lVert g\rVert_2=\sqrt{\sum_r\sum_{i\in r}g_i^2}$
因此一个简单的数学操作也会引入集合通信。
10. 把训练稳定性串起来
可以把本章的方法理解为针对不同阶段的防线:
1
2
3
4
5
6
mini-batch 用可承受的成本估计梯度
Momentum / Adam 让更新对噪声和尺度差异更稳健
初始化 让训练一开始的信号尺度正常
残差连接 给深层网络保留直接的梯度路径
LayerNorm/RMSNorm 在训练中持续控制隐藏状态尺度
梯度裁剪 限制偶发异常梯度造成的单步破坏
它们不能互相完全替代。一个现实的 Transformer 训练配置通常会同时使用残差、Norm、合适初始化、AdamW、学习率调度和梯度裁剪;每一项都在处理训练链路中不同位置的风险。