FP16数值上限65504易被突破导致inf/NaN,主因是attention中q@k.t()、exp(x)等操作溢出,AMP不拦截前向溢出,需手动插入.isfinite()检查关键张量。

PyTorch启用FP16(即torch.cuda.amp混合精度)后出现NaN,根本原因不是“训练崩了”,而是float16的数值表示能力太弱——它连65504都装不下,而你的模型中间结果轻轻松松就超了。
FP16的65504上限是怎么被踩爆的?
float16最大可表示正数是65504,不是近似值,是硬性天花板。一旦某个张量元素超过这个数,就会变成inf;后续任何含inf的运算(比如inf - inf、log(inf)、softmax里的exp(x))大概率产出NaN。
- 注意力层中
q @ k.t():当序列长×头维度较大时,点积结果极易突破65504,尤其在初始化权重偏大或输入未归一化时 -
exp(x)类操作:softmax前的logits若>11.1,exp(11.1)≈66700→ 直接溢出为inf - 损失函数输入非法:如
nn.CrossEntropyLoss接收了全零logits,或nn.BCEWithLogitsLoss的target含非0/1值,都会在内部计算中触发log(0)或除零
为什么AMP没拦住溢出?
AMP不是“全程护航”,它只对部分算子做自动类型升降级:conv2d、linear、matmul默认走FP16,但softmax、loss、optimizer.step()会回退到FP32。问题恰恰出在那些“被放行”的FP16算子上——它们的输出直接喂给下一个FP16算子,中间没有检查。
-
q @ k.t()用FP16算完是inf,紧接着inf / sqrt(d)还是inf,再进softmax(虽是FP32)也无力回天 -
GradScaler只缩放loss防止梯度下溢,对前向中的上溢(overflow)完全不干预 - FP16张量的
.isfinite()检查必须手动加,框架不会默认插入
怎么快速定位是哪一层在溢出?
别猜,用torch.autograd.detect_anomaly()配合逐层检查,重点盯三类张量:attention score、activation after exp、loss input。
立即学习“Python免费学习笔记(深入)”;
- 在
forward末尾插一句:assert torch.isfinite(x).all(), f"NaN/inf in {x.shape} at layer X" - 对
q @ k.t()结果立刻检查:if not torch.isfinite(attn_scores).all(): print(attn_scores.max()) - 避免全局
torch.autograd.set_detect_anomaly(True),它会让训练慢3倍以上,只在调试时临时开启
BatchNorm和LayerNorm为什么雪上加霜?
BN层在FP16下容易因统计量(尤其是方差)过小而失效:当var被舍入为0,1/sqrt(var + eps)就变成inf;LN同理,且其eps=1e-5在FP16里已接近最小可表示正数6.1e-5,防护作用大幅削弱。
- BN建议改用
nn.SyncBatchNorm或显式设track_running_stats=False避免统计量污染 - LN推荐把
eps从1e-5调大到1e-4或1e-3,确保在FP16下仍有意义 - 更稳妥的做法:对LN/BN的输入先做
.clamp(min=-10, max=10),主动限制动态范围
最易被忽略的一点:NaN不是孤立事件,而是“链式污染”。一个inf出现在第3层,可能到第12层才报NaN;所以检查点不能只设在loss处,得卡在每个高危算子之后。


















