PyTorch调用CrossEntropyLoss报张量布尔值歧义错误如何解决
报错原因
你触发这个错误的核心问题是混淆了PyTorch损失的类初始化和计算逻辑:
torch.nn.CrossEntropyLoss是损失函数类,直接在类名后加括号传参时,执行的是类的初始化构造方法,不是损失值计算逻辑。- 构造方法的第一个位置参数是自定义类别权重
weight,你把模型预测输出tensor r传给了这个参数,初始化流程中会对该参数做布尔合法性校验,包含多个元素的Tensor无法直接转换为单个布尔值,就抛出了对应的RuntimeError。
正确修改方案
两种标准写法都可以正常运行,根据自己的编码习惯选就行:
- 先实例化损失类,再传入预测值、标签计算损失
import torch l = torch.tensor([0, 1, 1, 1], requires_grad=False) r = torch.rand(4, 2) # 无自定义权重、忽略索引等特殊需求时,初始化参数留空即可 loss_fn = torch.nn.CrossEntropyLoss() # 前向计算传参:第一个位置是预测logits(形状要求[batch_size, 类别数]),第二个是类别标签 loss = loss_fn(r, l) print(loss)
- 直接调用函数式接口,无需提前实例化类
import torch import torch.nn.functional as F l = torch.tensor([0, 1, 1, 1], requires_grad=False) r = torch.rand(4, 2) loss = F.cross_entropy(r, l) print(loss)
额外注意点
- 传入CrossEntropyLoss的预测值不需要提前做softmax操作,损失内部会自动完成log_softmax和负对数似然计算,提前做softmax反而会导致结果异常。
- 索引格式的标签默认需要是
torch.long类型,如果你的标签tensor是浮点类型,可以提前加.long()转换,避免后续触发类型校验错误。
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

