在第 4 章前面几节中,我们讨论了因果语言建模(CLM)的训练目标、交叉熵损失的计算逻辑,以及 AdamW、Warmup 等优化策略。但当你真的把一套百亿参数规模的模型放进 GPU 集群开始训练时,最先遇到的往往不是“loss 下降太慢”,而是“loss 突然变成 NaN”或者“训练了三天,loss 纹丝不动”。这就是训练稳定性问题。
对于大模型预训练而言,稳定性不是“调优项”,而是“生死线”。一次失控的梯度爆炸,意味着价值数百万美元的算力瞬间归零。本节从三个核心机制入手,讲清楚深层网络训练中的数值稳定性本质,以及工业界如何在 FP16/BF16 的刀尖上跳舞。
一、梯度消失与梯度爆炸:反向传播的两面诅咒
要理解这两个问题,必须先回到反向传播(Backpropagation)的数学本质:链式法则。神经网络越深,反向传播时梯度就要经过越多的矩阵乘法和激活函数导数相乘。这意味着梯度信号在从输出层传回输入层的路上,要么不断衰减,要么指数级放大。
1. 梯度消失(Vanishing Gradient)
表现:靠近输入层的参数几乎不被更新,模型底层“冻住”,loss 卡在高位平台长期不动。
在大模型中的根源:
- 虽然 Transformer 已经用 ReLU/GELU 和注意力机制缓解了早期 RNN 中 sigmoid 带来的梯度消失(回忆 1.2 节 RNN 的退场),但在极深层网络(如 96 层的 GPT-3 级模型)中,梯度仍可能因连乘效应而衰减。
- 注意力机制中的 softmax 归一化本身会压缩梯度幅度;如果加上 LayerNorm 的前置或后置位置不当,某些层的梯度信号可能被过度抑制。
直观类比:梯度像是从山顶传向山脚的接力口令,每层传话时都小声一点,传到山脚时已经听不见了。
2. 梯度爆炸(Gradient Explosion)
表现:loss 在某个 step 突然飙升,几个 step 后变成 Inf/NaN;参数值在一瞬间膨胀到天文数字,模型彻底“脑死亡”。
在大模型中的根源:
- 大模型的注意力计算中包含 $QK^T/\sqrt{d_k}$ 这样的矩阵乘法(回忆 3.3 节)。在深层堆叠下,多个大维度的矩阵连乘极易导致数值溢出。
- 当 batch 内出现异常长序列或离群数据点时,梯度范数可能在单个 step 内暴涨几个数量级。
直观类比:接力口令越传越响,传到山脚时已经变成震耳欲聋的噪音,所有人都被震懵了。
二、稳定训练的四道工程防线
单纯依靠优化器本身无法驯服这些数值问题。现代 LLM 的训练稳定性是架构设计与工程技巧共同作用的结果。
1. 残差连接(Residual Connections):梯度高速公路
这是 3.6 节提到的核心组件,但值得从稳定性角度再强调。残差连接让每一层的输入 $x$ 直接叠加到输出 $F(x)+x$ 上。在反向传播时,梯度可以沿着这条“捷径”直接回流,而不必每次都经过 $F(x)$ 的复杂变换。这从根本上缓解了深层网络梯度消失的风险。
2. 层归一化(Layer Normalization):稳定分布
LayerNorm 将每一层的输入/输出强制归一到均值为 0、方差为 1 的分布。这避免了“Internal Covariate Shift”——即前面层的参数更新导致后面层输入分布剧烈变化的问题。没有 LayerNorm,深层 Transformer 的训练 loss 会在早期就发散。
实用细节:GPT-3 及之后的多数模型采用 Pre-LN(将 LayerNorm 放在残差分支之前),相比原始的 Post-LN,能让训练在更大学习率下保持稳定。
3. 梯度裁剪(Gradient Clipping):最后的保险杠
这是最朴实无华也最有效的救命手段。设定一个全局梯度的 L2 范数上限(如 max_norm=1.0),如果某个 step 计算出的梯度范数超过阈值,就将所有梯度按比例缩放回阈值内。
# 伪代码
global_norm = compute_gradient_norm(all_params)
if global_norm > max_norm:
scale = max_norm / global_norm
for g in gradients:
g *= scale
工业真相:几乎所有百亿级以上模型的预训练都会开启梯度裁剪。它不会显著影响收敛,但能在数据波动或学习率试探期拦住 90% 的 NaN 事故。
4. 学习率 Warmup:启动期的“怠速热车”
呼应 4.4 节,Warmup 在训练最初几百到几千个 step 内,将学习率从 0 线性增加到预设峰值。这是因为模型初始参数是随机的,早期就使用大学习率会直接撕毁参数空间的有序结构。Warmup 给模型一个“缓冲带”,让梯度方向先稳定下来。
三、混合精度训练:在精度与效率之间走钢丝
当我们讨论 2.1 节“70B 模型全精度需要 140GB 显存”时,其实已经触及了一个工程现实:FP32(单精度浮点)太慢、太占显存,但 FP16(半精度)又太容易数值爆炸。 混合精度训练(Mixed Precision Training)就是解决这个矛盾的工业级方案。
1. 为什么需要混合精度?
- 显存压力:FP32 每个参数占 4 字节,而 FP16 占 2 字节。对于 AdamW 优化器,它需要存储参数、动量(momentum)、二阶矩(variance),三者都按 FP32 存的话,显存开销惊人。
- 计算速度:NVIDIA Tensor Core 对 FP16/BF16 的矩阵乘法有原生硬件加速,吞吐量可达 FP32 的 8 倍。
2. FP16 vs BF16:两种半精度的性格差异
| 特性 | FP16 | BF16(Brain Float 16) |
|------|------|------------------------|
| 指数位 | 5 bit | 8 bit(与 FP32 相同) |
| 尾数位 | 10 bit | 7 bit |
| 动态范围 | 很窄(最小正规格数 ~6.1e-5) | 与 FP32 基本一致 |
| 精度 | 高 | 低(但足够训练) |
| 硬件支持 | Pascal 以后 | Ampere(A100)以后 |
关键认知:
- FP16 的动态范围太小:梯度更新值很容易下溢(Underflow)变成 0,导致参数停止学习;或者上溢变成 Inf。
- BF16 是更稳妥的选择:它用精度换范围。对于 LLM 训练,参数和梯度通常不需要小数点后十几位的精度,但需要足够的范围来避免溢出。因此,A100/H100 时代的大模型训练,BF16 已成为默认选项。
3. 自动混合精度(AMP)的工作原理
工业界普遍采用 NVIDIA 的 Automatic Mixed Precision (AMP) 策略,其核心逻辑是“FP16 算,FP32 保”:
- 前向传播:权重以 FP16 参与矩阵乘法,利用 Tensor Core 加速;
- 损失计算:Loss 以 FP32 保留,避免半精度累加误差;
- 反向传播:梯度以 FP16 计算,但在更新前转换为 FP32;
- 优化器状态:AdamW 的动量和二阶矩必须用 FP32 存储,否则累积误差会直接导致训练发散;
- 参数更新:用 FP32 的优化器状态计算更新量,再写回 FP16 的权重。
4. Loss Scaling:对抗 FP16 的下溢
即使使用 AMP,FP16 的梯度仍可能小到被舍入为 0。解决方法是 Loss Scaling(损失缩放):
- 在反向传播前,先将 loss 乘以一个较大的 scale factor(如 $2^{16}=65536$);
- 反向传播后的梯度也会同比例放大,从而逃离 FP16 的“死亡区间”;
- 更新参数前,再将梯度除以同样的 scale factor,还原真实数值。
现代框架(PyTorch AMP、DeepSpeed)通常支持动态 Loss Scaling:自动监测梯度是否出现 Inf/NaN,如果出现就下调 scale,如果稳定就逐步上调。
四、训练不稳定的排查清单:一个工程师的实战手册
当你面对训练日志时,以下是按优先级排序的排查逻辑:
| 症状 | 最可能原因 | 急救措施 |
|------|-----------|----------|
| Loss 在最初几个 step 后突变为 NaN | 学习率过大 / 初始化冲突 / 数据异常 | 降低 Warmup 峰值、检查输入序列长度、开启梯度裁剪 |
| Loss 长期平台期,几乎不动 | 梯度消失 / 学习率过小 / 数据质量差 | 检查 LayerNorm 位置、增大学习率、确认数据未经过度清洗 |
| Loss 震荡剧烈但不发散 | 学习率偏高 / batch size 偏小 | 适当降低学习率或增大梯度累积步数 |
| 混合精度下频繁 NaN,FP32 正常 | Loss Scaling 失败 / BF16/FP16 溢出 | 换用 BF16(如果硬件支持);若必须用 FP16,降低动态 loss scale 上限 |
| 特定 step 必现 NaN | 脏数据(超长序列、非法字符、空样本) | 加入异常数据过滤与长度截断 |
一条铁律:在启动正式的大规模训练前,务必先用 小模型(如把层数砍半)和 短序列(如 512 tokens)在相同数据上跑通完整 pipeline。如果在小尺度上都稳定不了,放大到千亿参数只会死得更快、更贵。
五、小结
梯度消失与梯度爆炸是深层网络反向传播的固有数学风险;残差连接、LayerNorm、梯度裁剪和 Warmup 共同构成了现代 LLM 训练的“稳定四件套”。在此之上,混合精度训练通过 FP16/BF16 与 FP32 的精密协作,将显存占用压缩近半、计算吞吐量翻倍,使大模型预训练在经济上成为可能。
关键 takeaway:
- 梯度裁剪是底线保险,无论你的模型多大、数据多干净,都应该开启;
- BF16 在绝大多数场景优于 FP16,前提是算力层支持 Ampere 及以上架构(A100、H100、RTX 30/40 系);
- 训练稳定性是系统工程,架构设计(Pre-LN)、优化策略(Warmup)、数值计算(AMP)和脏数据过滤,缺一不可。
至此,第 4 章关于预训练原理的讨论告一段落。从训练目标(CLM)、损失函数、优化器到本节的核心稳定性保障,我们已经完整覆盖了“如何让一个空白模型吃下海量文本并收敛为有序参数”的全过程。接下来在第 5 章中,我们将进入模型的“社会化”阶段——对齐技术:如何通过 SFT、RLHF 和 DPO,让这个会接龙的概率机器,变成一个听得懂人话、守规矩、有实用价值的对话助手。