训练发散是深度学习中最令人头疼的问题之一:上一轮还在正常收敛的模型,这一轮loss突然飙到几十万,或者直接变成NaN,整个训练宣告失败。造成发散的原因五花八门,但统计下来,绝大多数案例都指向两个环节:参数初始化不当和梯度在反向传播中失去控制。这篇文章围绕这两点展开,把原理讲清楚,再给出可直接落地的代码方案。

训练发散的根源:数值为什么会失控
要理解发散,先要看清数值在网络中是如何流动的。前向传播时,每一层的输出是上一层输出乘以权重再经过激活函数。如果权重初始值偏大,每经过一层,数值方差会被放大一次,经过几十层之后,激活值轻松溢出浮点数上限,loss直接变成inf或NaN。反过来,如果权重初始值偏小,信号逐层衰减,到达浅层时梯度几乎为零,网络学不到东西,这叫梯度消失。
反向传播遵循同样的逻辑,只不过方向相反。根据链式法则,梯度是各层局部导数的连乘积。当每一项都大于1时,连乘结果指数级增长,这就是梯度爆炸;当每一项都小于1时,连乘结果指数级衰减,这就是梯度消失。梯度爆炸的直接后果是优化器一步跨得太远,参数被甩到很远的区域,loss不降反升,最终崩溃。
判别方法很简单:在训练循环里打印梯度范数。如果发现某些step上梯度范数从正常的个位数突然跳到上千甚至上万,基本可以确认是梯度爆炸。而如果loss变成NaN且梯度范数早在此之前就已经异常增大,说明爆炸早就发生了,只是浮点溢出滞后了一步。
for name, param in model.named_parameters():
if param.grad is not None:
print(name, param.grad.norm().item())参数初始化:稳住前向传播的第一道防线
初始化的目标是让信号在通过网络时不被逐层放大或缩小。数学上的表述是:保持每层输出的方差与输入方差大致相等。Xavier初始化(也叫Glorot初始化)就是为这个目标设计的,它把权重采样在一个与输入输出维度相关的区间内,使得方差在前向和反向两个方向上都保持稳定。它最适合搭配tanh、sigmoid这类对称激活函数。
但Xavier初始化有一个隐含假设:激活函数在零点附近近似线性。ReLU函数负半轴恒为零,会“砍掉”约一半的激活值,方差天然减半,Xavier的假设被打破。Kaiming初始化针对这一点做了修正,在方差计算中补上了因子2,专门服务于ReLU家族。一个常见的错误是:用PyTorch默认初始化搭了一个很深的ReLU网络,却不手动调整,结果深层的激活方差逐层衰减,训练慢得像蜗牛,还以为是学习率设小了。
import torch.nn as nn
class DeepNet(nn.Module):
def __init__(self, depth=20, width=256):
super().__init__()
layers = []
for _ in range(depth - 1):
layers.append(nn.Linear(width, width))
layers.append(nn.ReLU())
layers.append(nn.Linear(width, 10))
self.net = nn.Sequential(*layers)
self._init_weights()
def _init_weights(self):
# 对Linear层使用Kaiming初始化,适配ReLU激活
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.kaiming_normal_(m.weight, nonlinearity='relu')
if m.bias is not None:
nn.init.zeros_(m.bias)几点实践经验值得强调。第一,BatchNorm或LayerNorm会在一定程度上掩盖初始化问题,因为归一化层强行把分布拉回来了,所以加了归一化的网络对初始化没那么敏感,但残差连接的分支输出仍需小心。第二,bias通常初始化为零即可,不需要花哨的处理。第三,若使用了LeakyReLU,Kaiming公式中的因子需要按负斜率微调,不过实践中影响不大,默认配置往往够用。
梯度裁剪:给反向传播装上安全阀
即便初始化没问题,训练中期仍可能出现梯度爆炸,尤其是RNN、LSTM这类循环结构,以及使用了较大学习率的场景。梯度裁剪的思路很直接:当梯度范数超过阈值时,按比例缩小梯度,相当于给参数更新步长设置上限。它不改变梯度方向,只限制幅度,因此对正常收敛几乎无干扰。
按范数裁剪是最常用的方式。它先计算所有参数梯度的整体L2范数,若超过阈值max_norm,就将梯度整体乘以系数max_norm除以当前范数。PyTorch封装得非常成熟,一行代码即可:
optimizer.zero_grad() loss.backward() # 全局范数裁剪,阈值通常在0.5到5之间试探 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step()
另一种是按值裁剪,把每个梯度元素强制限制在区间内。这种方式实现简单,但改变了梯度的方向分布,一般只用于训练RNN时做保守兜底,或者排查问题的临时手段。阈值的选择没有万能值,常见做法是先观察正常训练阶段梯度范数的分布,取其若干倍作为阈值,既保证平时裁剪不生效,又能在异常时及时刹车。
# 按值裁剪:把每个梯度元素截断到[-0.5, 0.5] torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
需要澄清一个误区:梯度裁剪治标不治本。如果梯度爆炸频繁触发裁剪,说明模型本身存在问题,可能是数据里有异常样本、学习率过大或者损失函数实现有bug(比如手写的log操作没有加epsilon保护)。裁剪应该被视为保险丝,而不是常规组件。训练Transformer时几乎默认开启裁剪,是因为注意力机制在早期训练阶段确实容易不稳定,属于合理的工程妥协。
排查训练发散的完整流程
当发散真的发生时,按固定顺序排查能节省大量时间。第一步,定位发散位置:回溯打印loss、梯度范数和部分激活值的统计量,确认是前向溢出还是反向爆炸。第二步,检查数据:训练集中混入的NaN、inf或量纲异常的样本是常见元凶,用一行检测代码可以快速扫出问题样本。第三步,检查损失函数:log之前加1e-8,除法之前加epsilon,开方之前做clamp,这些防御性写法能挡住大部分数值陷阱。
import torch
# 检测批次数据中是否混入异常值
def check_batch(batch):
if isinstance(batch, torch.Tensor):
assert not torch.isnan(batch).any(), "数据中出现NaN"
assert not torch.isinf(batch).any(), "数据中出现inf"
print("min:", batch.min().item(), "max:", batch.max().item())
# 防御性的交叉熵不需要手动处理,但自定义损失要小心
def safe_log(x, eps=1e-8):
return torch.log(torch.clamp(x, min=eps))第四步,降低学习率做对照实验。把学习率除以10再训练,如果发散消失,说明步长确实过大,可以换用warmup策略,让学习率从很小的值线性爬升,避开训练初期最脆弱的阶段。第五步,确认混合精度训练的配置。使用AMP时,损失缩放系数设置不当也会导致梯度下溢或上溢,建议开启GradScaler并配合动态调整。
总结一套稳妥的默认配置:深层网络使用Kaiming或Xavier初始化并在首个batch后检查激活方差;优化器更新前调用按范数裁剪,max_norm取1.0;学习率配合warmup;自定义损失函数全部加数值保护。这四件事成本极低,却能过滤掉绝大多数训练发散问题。剩下的顽固案例,往往隐藏在数据或模型结构的细节里,这时候冷静地打印中间变量,比盲目调参有效得多。