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

PyTorch自定义指数交叉熵损失函数训练时突发nan问题求助

核心触发原因
  • 指数运算数值溢出:你自定义的损失用到了torch.exp,而交叉熵损失的取值上限本身没有边界:当模型对某个batch的样本预测完全错误时,交叉熵损失会趋近于正无穷,除以η之后输入指数函数,很容易超出PyTorch浮点数的表示范围,直接得到inf值,后续反向传播、参数更新时就会连锁产生nan。
  • 梯度爆炸:该损失的梯度为 $\frac{1}{\eta} \times exp(L/\eta) \times \nabla L$,本身exp的输出量级已经很高,梯度也会随之爆炸,超出数值范围后变成nan,直接冲毁模型参数,导致后续所有迭代的输出异常。
  • 补充:你当前的平均损失计算逻辑有误,累加时乘以了batch样本数data.size(0),最后却只除以了batch数len(train_loader),得到的不是单样本平均损失,但这个不是触发nan的核心原因。
修复方案
  • 损失值截断:计算指数后增加数值截断逻辑,避免出现inf,修改损失行代码为:
loss = torch.exp(criterion(output, target)/eta).clamp(min=1e-8, max=1e6)
  • 增加梯度裁剪:在反向传播之后、优化器更新之前,添加梯度裁剪代码,避免梯度冲毁参数:
loss.backward()
torch.nn.utils.clip_grad_norm_(client_model.parameters(), max_norm=1.0) # 新增行
optimizer.step()
  • 调整η取值:如果当前η取值过小,可以适当调大η,降低指数函数输入的量级,减少溢出概率。
  • 异常batch过滤:计算损失后可以增加异常判断,遇到inf/nan损失时跳过当前batch的更新,避免污染模型参数:
loss = torch.exp(criterion(output, target)/eta)
if torch.isnan(loss) or torch.isinf(loss):
    continue
loss.backward()
  • 修正平均损失计算:将平均损失的除数改为训练集总样本数:
train_loss = train_loss/len(train_loader.dataset)

内容的提问来源于stack exchange,提问作者Chao

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 13:45:09