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
相关产品推荐
相关产品推荐

