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
相关产品推荐
相关产品推荐

