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

PyTorch多分类逻辑回归训练报错:请求排查并修正代码问题

问题分析与修正方案

错误点总结

  1. 标签张量格式与类型错误
    CrossEntropyLoss要求目标标签是1D整数张量(类别索引),但你的代码中把y处理成了2D的float张量(y = torch.tensor(y).float().unsqueeze(1)),直接触发了"multi-target not supported"错误。
  2. 模型输出激活函数误用
    多分类场景下搭配CrossEntropyLoss时,模型最后一层不需要加sigmoid激活——CrossEntropyLoss内部已经集成了log_softmax计算逻辑,额外的sigmoid会破坏损失计算的正确性。
  3. 模型实例化顺序错误
    你先实例化了model = LR(n_features, n_outputs),但n_features和n_outputs是在后续代码中才定义的,这会导致运行时出现变量未定义的报错。

具体修正步骤

1. 修正标签处理逻辑

在数据加载函数中,将标签改为1D的long类型张量,移除不必要的维度扩展和类型转换:

y = df['label']
y = torch.tensor(y).long()  # 转为long类型的1D张量,匹配CrossEntropyLoss要求

2. 修正模型的forward方法

移除sigmoid激活,直接返回线性层的原始输出:

class LR(torch.nn.Module):
    def __init__(self, n_features, n_outputs):
        super(LR, self).__init__()
        self.lr = torch.nn.Linear(n_features, n_outputs)
        
    def forward(self, x):
        out = self.lr(x)  # 去掉sigmoid激活
        return out

3. 调整模型实例化顺序

先定义模型所需的参数,再实例化模型:

# 先定义参数
n_features = 384 
n_outputs = 7
# 再实例化模型及相关组件
model = LR(n_features, n_outputs)
optim = torch.optim.SGD(model.parameters(), lr=0.01)
criterion = torch.nn.CrossEntropyLoss()

修正后的完整训练代码片段

class LR(torch.nn.Module):
    def __init__(self, n_features, n_outputs):
        super(LR, self).__init__()
        self.lr = torch.nn.Linear(n_features, n_outputs)
        
    def forward(self, x):
        out = self.lr(x)
        return out

# 先定义参数再实例化模型
n_features = 384 
n_outputs = 7
model = LR(n_features, n_outputs)
optim = torch.optim.SGD(model.parameters(), lr=0.01)
criterion = torch.nn.CrossEntropyLoss()

EPOCHS = 6

def train(model, optim, criterion, x, y, epochs=EPOCHS):
    for e in range(1, epochs + 1):
        optim.zero_grad()
        out = model(x)
        loss = criterion(out, y)
        loss.backward()
        optim.step()
        print(f"Loss at epoch {e}: {loss.data}")
    return model

model = train(model, optim, criterion, x_train, y_train)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 05:24:54