在Transformer的Attention计算中,Softmax函数负责对查询和键的点积分数做归一化。当序列长度变大或分数绝对值偏高时,指数运算会让数值迅速膨胀,最终超出浮点数表示范围,出现Inf或NaN,这就是典型的Attention计算溢出问题。解决它的核心思路是在不改动数学等价性的前提下,提升Softmax的数值稳定性。

为什么Softmax会发生数值溢出
标准Softmax的定义为对向量x的每个元素计算exp(x_i)再除以所有exp的和。假设x里有一个值是1000,那么exp(1000)在单精度浮点数下远远超过了约3.4e38的上限,直接变成Inf。一旦分子或分母出现Inf,后续除法和梯度回传都会失效,训练被迫停止。
在Attention里,分数矩阵由Q和K相乘并除以根号d得到。如果维度较大或没有合理缩放,某些分值可能非常极端。尤其是批量训练时,不同样本的最大值差异巨大,不加处理就会频繁触发溢出。理解这一点,才能明白稳定性优化并不是可选项,而是工程实现的必选项。
减去最大值:最简单有效的等价变换
数学上,给x的每个元素同时减去常数c,Softmax结果不变,因为分子分母的exp(c)可以约掉。通常取c为x里的最大值max(x),这样平移后的最大指数为exp(0)=1,其余均为负数指数,不会超过1,从根源上杜绝了上溢。
举例来说,若分数为[1000, 1001, 999],直接算会溢出;减去最大值1001后变成[-1, 0, -2],exp结果分别是[0.3679, 1, 0.1353],求和后再归一化,得到稳定概率。这种操作在PyTorch的softmax里已默认集成,但手写Attention时必须显式写上,否则自定义核函数仍会出错。
下溢与log_sum_exp的配合
减去最大值解决了上溢,但小数指数可能下溢成0,导致求和后为0进而除零。为此在算对数似然或交叉熵时,应使用log_sum_exp技巧:先取最大a,再算log(sum(exp(x-a)))+a,以对数域安全求和。
例如计算log_softmax,可写为x - (x.max() + log(sum(exp(x - x.max()))))。这样即使个别项下溢为0,求和项依然保留主要贡献,不会让分母为0,也方便后续与标签做稳定交叉熵。
Attention中的具体优化实现
在自注意力分数S上,应先沿最后一个维度求最大值,再广播减去。伪代码逻辑为:S = S - S.max(dim=-1, keepdim=True);然后e = exp(S);P = e / e.sum(dim=-1, keepdim=True)。这样无论序列多长,概率矩阵始终有限。
如果还做mask,要把padding位置设成负无穷,减去最大值后它们变成负无穷减有限值仍为负无穷,exp后为0,求和不受影响。但要注意某些框架负无穷减负无穷会得到NaN,因此mask应在减最大值前用大负数(如-1e9)替代,而不是真正inf。
缩放与混合精度下的额外注意
除以根号d是一种软性稳定,降低分值方差。在混合精度用float16时,最大正值仅约65504,所以减最大值更为关键。建议先在float32算Softmax再转回半精度,或确保减最大值后指数和不超过半精度上限。
另外,一些高效Attention库采用分块计算最大值再归一,逻辑相同但需注意跨块最大值同步。只要遵循先减最大、后指数、再归一的顺序,就能在绝大多数硬件上避免溢出。
常见错误与排查清单
不少实现漏写keepdim,导致广播形状错误,概率和不为1;还有人在mask时使用bool直接乘0,却忘了负无穷已污染最大值。下表列出典型问题和对策。
| 问题现象 | 可能原因 | 修复方式 |
|---|---|---|
| Loss变成NaN | 未减最大值直接exp | 显式减去行最大值 |
| 概率和不为1 | keepdim缺失 | 保持维度一致再除和 |
| mask位置概率非0 | 用inf后减inf得NaN | mask用-1e9代替inf |
只要按上述方式组织代码,Attention里的Softmax数值稳定性优化就能落地,长文本训练也不再被溢出中断。数值稳定不是理论细节,而是保障模型可训的底线。
Softmax数值稳定性Attention计算溢出log_sum_exp修改时间:2026-08-11 20:39:26