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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 09:45:39