You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.09 07:32:28