1. FP16训练中的数值稳定性危机
在深度学习训练中,FP16(半精度浮点数)因其内存占用小、计算速度快等优势,已成为现代训练流水线的标配。但鲜为人知的是,这个看似高效的"加速器"背后,隐藏着一个足以毁掉整个训练过程的隐形杀手——数值溢出导致的NaN(Not a Number)污染。
FP16的数值表示范围极其有限,仅能表示-65504到65504之间的数字。当数值超过这个范围时,就会发生上溢(overflow)或下溢(underflow)。在softmax计算中,指数运算$e^x$会迅速放大输入值,$e^{12}≈162754$就已经超过了FP16的最大表示范围,计算结果会变成INF(无穷大)。而一旦出现INF,后续的除法操作就会产生NaN,这种污染会像病毒一样在整个计算图中传播。
关键事实:在BERT-large等现代神经网络中,中间层的激活值超过12的情况相当普遍。这意味着如果不加处理,直接使用FP16计算softmax,NaN的出现不是"是否"的问题,而是"何时"的问题。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. Softmax的数值稳定化原理
2.1 数学基础:平移不变性
Softmax函数具有一个关键数学性质——平移不变性:
$$
\text{Softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}} = \frac{e^{x_i - M}}{\sum_j e^{x_j - M}}
$$
其中$M$可以是任意实数,通常取$M=\max(x)$。这一性质让我们可以对输入进行平移而不改变输出结果。
通过减去最大值$M$,我们确保所有指数运算的输入$x_i - M$都≤0,因此$e^{x_i - M}$的范围被限制在(0,1]之间,彻底避免了上溢风险。虽然极小的值仍可能下溢为0,但这通常不会导致NaN,只是损失一些精度。
2.2 实现中的双重挑战
在实际实现中,我们需要同时解决两个问题:
- 数值稳定性:确保计算过程不会产生INF或NaN
- 计算效率:尽量减少数据搬运和类型转换的开销
一个典型的错误实现如下:
python复制def unsafe_softmax(x):
exp_x = np.exp(x) # 可能溢出
return exp_x /
