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

PyTorch使用Omniglot数据集报RuntimeError: Found dtype Long but expected Float解决

代码修复方案

你的代码存在两处核心错误,按如下步骤修改即可解决报错:

1. 全连接层输入维度计算错误

错误原因

卷积层输出展平后的维度计算错误,导致模型输出形状和标签形状不匹配:
逐层计算特征图尺寸:

  • 输入:[batch, 1, 105, 105]
  • 经过conv1+maxpool:输出[batch, 8, 51, 51]
  • 经过conv2+maxpool:输出[batch, 16, 24, 24]
  • 展平后维度为16 * 24 * 24 = 9216,不是代码中写的18432

修复代码

# 修改模型__init__方法中的fc1定义
self.fc1 = torch.nn.Linear(9216, 256)
# 修改forward方法中的展平操作
hidden = hidden.view(-1, 9216)

2. 损失函数选择错误

错误原因

分类任务错误使用了回归任务的MSELoss:

  • MSELoss要求输入和目标均为浮点类型、形状完全匹配,而你的标签是代表类别的长整型,形状仅为[batch_size],和模型输出[batch_size, 类别数]不匹配,触发类型错误
  • 分类任务应使用CrossEntropyLoss,原生支持分类任务的输入输出格式,无需额外转换数据类型

修复代码

# 替换损失函数定义
loss_fun = torch.nn.CrossEntropyLoss()

可选优化建议

  • 调整输出层维度:Omniglot训练集共964个类别,你当前设置的50类不符合数据集实际情况,根据任务需求修改fc3的输出维度:
    self.fc3 = torch.nn.Linear(64, 964) # 全量分类场景
    
  • 修正批次计数:直接用len(train_generator)获取总批次,避免整除误差:
    num_of_batch = len(train_generator)
    
  • 优化准确率计算:用浮点除法得到更精确的结果:
    print('Validation accuracy: %.2f%%' % ((100 * correct) / total))
    

内容的提问来源于stack exchange,提问作者Seng

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 02:57:03