运行PyTorch脚本遇类型错误求助:Expected Long but got Int
解决PyTorch中"target期望Long类型但得到Int"的错误
嗨,我来帮你搞定这个PyTorch的类型不匹配问题!这个错误大多出现在分类任务里,比如用CrossEntropyLoss这类损失函数的时候——PyTorch明确要求目标标签(也就是参数里的target)得是torch.long(64位整数)类型,但你的数据现在是torch.int(32位整数),类型不兼容就触发了这个报错。
给你几个实用的解决办法,按需选就行:
方法一:提前转换目标数据类型
如果你的target是张量,直接调用target = target.long()就能完成转换。要是是从numpy数组转成张量的话,转的时候直接指定dtype更省心:target_tensor = torch.from_numpy(target_np).long()方法二:计算损失时临时转换
要是不想改动原始数据,在调用损失函数的时候临时转换也可以,比如写成:loss = criterion(outputs, target.long())这样每次计算损失时都会自动把target转成正确的类型。
方法三:在数据加载环节固定类型
如果你用了Dataset和DataLoader来加载数据,可以在Dataset的__getitem__方法里直接把target转成long类型,一劳永逸:def __getitem__(self, idx): data = self.data[idx] target = self.targets[idx] # 这里把target指定为long类型 return torch.tensor(data), torch.tensor(target, dtype=torch.long)
最后提个小提醒:用CrossEntropyLoss的时候,千万别给target做one-hot编码,直接传类别索引就好,而且这个索引必须是long类型,这可是很多新手容易踩的坑哦!
内容的提问来源于stack exchange,提问作者karthik reddy
相关产品推荐
相关产品推荐

