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

PyTorch实现逻辑回归时损失函数输出NaN的问题及解决方法

解决PyTorch逻辑回归损失输出NaN的问题

核心问题分析

你的代码出现NaN主要有三个原因:

  • 手动实现二元交叉熵时,sigmoid输出趋近0或1时,th.log(0)会产生负无穷,后续运算触发NaN。
  • Adam优化器学习率设置过高(0.1),参数更新幅度过大导致模型发散。
  • 数据维度存在不匹配,可能引发广播计算错误。

具体修复方案

1. 替换手动损失函数为内置稳定实现

PyTorch的BCEWithLogitsLoss将sigmoid激活和二元交叉熵计算合并,通过数值稳定的方式避免log(0)问题,是二元分类的最优选择。

2. 降低学习率

Adam默认学习率为0.001,0.1的学习率会导致参数震荡甚至发散,调整为0.001或更小。

3. 修正数据维度匹配

确保模型输出和标签的形状完全一致,避免广播错误。

修正后的完整代码

数据加载与处理

from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
import torch as th
import numpy as np

dataLoad = load_breast_cancer()
X_train, X_test, Y_train, Y_test = train_test_split(dataLoad.data, dataLoad.target, test_size=0.33)
scaler = StandardScaler()
X_train = scaler.fit_transform(X_train)
X_test = scaler.transform(X_test)

# 调整为PyTorch常规输入格式:(样本数, 特征数)
X_train = th.from_numpy(X_train.astype(np.float32))
# 将标签转为二维张量,与模型输出形状匹配
Y_train = th.from_numpy(Y_train.astype(np.float32)).unsqueeze(1)

模型与损失函数定义

# 用PyTorch内置模块定义模型,更规范易维护
class LogisticRegression(th.nn.Module):
    def __init__(self, input_dim):
        super().__init__()
        self.linear = th.nn.Linear(input_dim, 1)
    
    def forward(self, x):
        # 直接返回logits,交给BCEWithLogitsLoss处理sigmoid和损失计算
        return self.linear(x)

# 使用内置数值稳定的损失函数
loss_fn = th.nn.BCEWithLogitsLoss()

训练循环

input_dim = X_train.shape[1]
model = LogisticRegression(input_dim)
optimizer = th.optim.Adam(model.parameters(), lr=0.001)  # 降低学习率至合理范围
epoch = 1000
losses = []

for i in range(epoch):
    optimizer.zero_grad()
    logits = model(X_train)
    loss = loss_fn(logits, Y_train)
    loss.backward()
    optimizer.step()
    losses.append(loss.item())
    
    # 可选:每100轮打印损失,监控训练状态
    if (i+1) % 100 == 0:
        print(f"Epoch {i+1}, Loss: {loss.item():.4f}")

额外说明

如果坚持要手动实现损失函数,需要给sigmoid输出添加极小值约束,避免log(0):

def lossFunction(predictY, outputY):
    epsilon = 1e-7
    # 将输出限制在[epsilon, 1-epsilon]范围内,避免log(0)
    predictY = th.clamp(predictY, epsilon, 1 - epsilon)
    loss = -(outputY * th.log(predictY) + (1 - outputY) * th.log(1 - predictY)).mean()
    return loss

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 06:10:38