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

PyTorch中torch.nn.Functional多分类整数标签交叉熵损失问题

问题解答

问题1:torch.nn.Functional中是否存在可计算多分类(每个实例对应一个整数标签)交叉熵损失的方法?

有,F.cross_entropy本身就支持整数标签的多分类交叉熵计算,但你用错了输入格式:

  • F.cross_entropy的第一个参数input需要是模型输出的logits(未经过softmax的原始得分),形状必须是[batch_size, num_classes],数据类型为浮点型;
  • 第二个参数target是整数形式的真实标签,数据类型为LongTensor,形状为[batch_size]即可。

你当前传入的predictions是直接的类别索引(比如1、2这类),完全不符合input的要求,这才是报错的核心原因,并非标签类型问题。

问题2:是否需要将两个列表转换为FloatTensor?

不需要,两者的类型要求正好相反:

  • 预测值(logits):必须转为浮点型张量(如FloatTensor),且形状要符合[batch_size, num_classes];
  • 真实标签:需要转为LongTensor(整数类型),绝对不能转成FloatTensor。

修正后的示例代码

import torch
import torch.nn.functional as F

# 假设总共有8个类别(因为真实标签最大值是7)
num_classes = 8

# 实际场景中,logits是模型的直接输出,这里用随机张量模拟
logits = torch.randn(5, num_classes)  # 形状[5,8],浮点型,对应5个样本、8个类别

actual_targets = [1, 2, 6, 5, 7]
targets = torch.tensor(actual_targets, dtype=torch.long)  # 转为LongTensor

# 正确计算交叉熵损失
loss = F.cross_entropy(logits, targets)
print(loss)

报错原因解析

你传入F.cross_entropy的predictions是形状为[5]的LongTensor(类别索引),PyTorch会误判你在传入类别概率分布(这种场景下要求target是浮点型概率),因此抛出"Expected floating point type for target with class probabilities, got Long"的错误。本质是预测值的格式完全错误,和标签类型无关。

另外,你之前的类型转换代码存在冗余:predictions.type(torch.LongTensor)不会修改原张量,需要重新赋值(比如predictions = predictions.type(torch.LongTensor)),但这不是解决问题的关键。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 11:02:45