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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 21:09:03