PyTorch搭建CNN报错:输入设为Long类型却出现类型不匹配错误
解决PyTorch Conv2d的类型不匹配错误
错误原因
PyTorch的nn.Conv2d层默认使用浮点型权重(torch.FloatTensor或torch.cuda.FloatTensor),而你传入的输入是torch.LongTensor。卷积运算要求输入张量与层权重的数值类型完全一致,类型不匹配就会触发RuntimeError报错。
解决方法
核心是将输入张量转换为浮点型,有两种常用方式:
方式1:创建数据时直接转换
在生成数据后调用.float()方法转为浮点型,同时建议将图像数据归一化到0-1区间(CNN训练的常规操作):data = torch.randint(low=0, high=255, size=[2, 1, 1024, 1024], dtype=torch.int64).float() / 255.0方式2:传入模型前临时转换
如果需要保留原始长整型数据,在传入模型时转换:out = model(data.float())
修正后的完整代码
import torch import torch.nn as nn data = torch.randint(low=0, high=255, size=[2, 1, 1024, 1024], dtype=torch.int64).float() / 255.0 model = nn.Conv2d(1, 3, kernel_size=3, padding=1, bias=False) print(data.type()) # 输出torch.FloatTensor out = model(data) print(out.shape) # 输出torch.Size([2, 3, 1024, 1024])
注意:不要尝试将模型权重转为长整型,卷积运算本质是浮点运算,强制转换会丢失精度,完全不符合CNN的设计逻辑。
内容的提问来源于stack exchange,提问作者The Hagen
相关产品推荐
相关产品推荐

