RESNet18处理float32单通道数据时target类型报错如何解决
错误原因
这个报错是PyTorch交叉熵损失的类型校验触发的:nn.CrossEntropyLoss要求传入的第二个参数target(也就是样本标签)必须是torch.long(int64)类型的类别索引值,但你当前数据加载返回的标签是float32类型,因此类型不匹配。
注意:你使用的float32类型的图像输入本身是符合ResNet18的输入要求的,不需要调整输入数据的类型,仅需要调整标签的类型即可。
修复方案
可以选择以下任意一种方式修复:
- 方案1:在训练步骤中临时转换标签类型(改动最小,适配性最强)
修改你代码中的training_step方法,拿到标签后先转换为long类型:
def training_step(self, batch, batch_no): # implement single training step x, y = batch # 新增标签类型转换逻辑 y = y.long() logits = self(x) loss = self.loss(logits, y) return loss
如果你的代码后续加了验证、测试步骤,对应步骤里的标签也要做同样的转换。
- 方案2:在数据集定义中统一转换标签类型
如果你使用自定义的PyTorch Dataset加载数据,可以在__getitem__方法返回标签的时候,直接将标签转为long类型,从数据源层面解决类型不匹配问题,不需要在每个训练/验证步骤重复写转换逻辑:
def __getitem__(self, index): # 原有加载图像、读取标签的逻辑 img = ... # float32类型的单通道图像 label = ... # 原始读取的float32类型标签 # 新增转换逻辑 label = label.long() return img, label
内容的提问来源于stack exchange,提问作者DLH
相关产品推荐
相关产品推荐

