PyTorch中Softmax输出出现NaN与负值的解决方法咨询
解决Softmax输出负概率与NaN问题:归一化Softmax的PyTorch实现
你遇到的这个问题其实挺常见的——理论上Softmax的输出应该是[0,1]区间的概率值且总和为1,但当输入的logits(也就是Softmax层的输入)数值过大时,很容易出现数值上溢/下溢,进而导致NaN或者看起来像“负概率”的异常值(本质是数值计算误差)。
先给你直接上数值稳定版Softmax的PyTorch实现,不管是作为函数还是可复用的Module都能用,而且前后向传播完全顺畅:
方案1:自定义稳定Softmax函数
def stable_softmax(logits, dim=-1): # 核心技巧:减去当前维度的最大值,避免指数运算溢出 max_logits = torch.max(logits, dim=dim, keepdim=True)[0] exp_logits = torch.exp(logits - max_logits) return exp_logits / torch.sum(exp_logits, dim=dim, keepdim=True)
这个实现利用了Softmax的数学性质:Softmax(x) = Softmax(x - c)(c为任意常数),减去max能把logits的最大值拉到0,这样指数运算后不会出现极大值导致溢出,从根源上避免NaN,同时保证输出结果和原始Softmax完全一致。
方案2:封装成Module(方便集成到模型)
如果要像PyTorch内置层一样使用,可以封装成nn.Module:
import torch.nn as nn import torch class StableSoftmax(nn.Module): def __init__(self, dim=-1): super().__init__() self.dim = dim def forward(self, logits): max_logits = torch.max(logits, dim=self.dim, keepdim=True)[0] exp_logits = torch.exp(logits - max_logits) return exp_logits / torch.sum(exp_logits, dim=self.dim, keepdim=True)
使用时直接替换原来的nn.Softmax即可:
# 原来的代码 # self.softmax = nn.Softmax(dim=-1) # 替换为 self.softmax = StableSoftmax(dim=-1)
额外的关键提醒
除了替换Softmax,还有个容易忽略的点可能是你的损失函数搭配问题:
如果你的损失函数用的是
nn.CrossEntropyLoss(),那模型末尾绝对不要加Softmax层!
因为CrossEntropyLoss内部已经集成了数值稳定的Softmax + Log似然计算,自己再加Softmax会导致双重计算,反而引发数值不稳定,这很可能是你出现NaN的元凶之一。
另外,虽然你已经用了clip_grad_norm_,但如果前面层的输出数值波动太大,也可以考虑在关键层后加入nn.LayerNorm来稳定数值分布,进一步减少溢出风险。
内容的提问来源于stack exchange,提问作者Granth
相关产品推荐
相关产品推荐

