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

PyTorch调用CrossEntropyLoss报张量布尔值歧义错误如何解决

报错原因

你触发这个错误的核心问题是混淆了PyTorch损失的类初始化和计算逻辑:

  • torch.nn.CrossEntropyLoss 是损失函数类,直接在类名后加括号传参时,执行的是类的初始化构造方法,不是损失值计算逻辑。
  • 构造方法的第一个位置参数是自定义类别权重weight,你把模型预测输出tensor r传给了这个参数,初始化流程中会对该参数做布尔合法性校验,包含多个元素的Tensor无法直接转换为单个布尔值,就抛出了对应的RuntimeError。
正确修改方案

两种标准写法都可以正常运行,根据自己的编码习惯选就行:

  1. 先实例化损失类,再传入预测值、标签计算损失
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)
  1. 直接调用函数式接口,无需提前实例化类
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:12:13