PyTorch使用nn.CrossEntropyLoss()出现Long/Float类型不匹配报错原因
PyTorch CrossEntropyLoss 标签类型报错原因解析
核心原因
nn.CrossEntropyLoss() 要求传入的分类标签张量必须为 Long(长整型) 类型,你遇到的报错本质是第一种方法生成的标签张量数据类型不符合要求。
报错信息:
expected scalar type Long but found Float in PyTorch
两种实现的本质差异
- 第一种实现逻辑:
torch.zeros(100)、torch.ones(100)这类API默认生成torch.float32类型的浮点张量,torch.hstack仅做张量拼接不会修改数据类型,最终得到的labels为浮点型,直接触发类型不匹配报错。 - 第二种实现逻辑:
你在numpy数组转PyTorch张量后主动调用了.long()方法,强制将数据类型转换为符合要求的长整型,后续的torch.squeeze仅调整张量形状、不会修改数据类型,因此可以正常传入损失函数计算。
验证方法
你可以分别打印两种实现得到的标签的数据类型,就能直观看到差异:
# 查看第一种实现的标签类型 print(labels.dtype) # 查看第二种实现的标签类型 print(labels.dtype)
第一种实现的修复方案
只需要在拼接完成后主动做长整型转换即可:
labels = torch.hstack((torch.zeros(100),torch.ones(100),1+torch.ones(100))).long()
内容的提问来源于stack exchange,提问作者Hamed Farahani
相关产品推荐
相关产品推荐

