PyTorch自定义conv2d滤波器因输入张量类型不匹配触发运行时错误
问题原因分析
- 报错中出现CUDA相关提示的原因:你调用的
torchsummary库默认会优先检测环境中是否有可用CUDA设备,若存在会自动将summary功能用到的测试输入张量放到CUDA上,哪怕你没有手动移动自己的业务数据集,这个测试输入也会是CUDA上的torch.cuda.FloatTensor类型,这就是报错里出现cuda字样的核心原因。 - 类型不匹配的原因:你自定义的卷积核
w是从np.ones生成的,numpy数组默认浮点数精度为float64,转成torch张量后对应torch.DoubleTensor(即float64类型),而且这个权重是在forward方法内临时生成的,没有和模型绑定,不会自动跟随输入/模型的设备、精度做转换,因此出现了输入(CUDA、float32)和权重(CPU、float64)的设备、类型都不匹配的问题。
修复方案
- 如果你不需要用CUDA跑模型验证,可在调用summary时指定运行设备为CPU即可解决CUDA相关的设备不匹配问题:
summary(model, (1,10,10), device='cpu')
- 建议将自定义卷积核注册为模型的缓冲区,这样模型切换设备、精度时权重会自动同步,同时避免每次forward重复生成权重,完整修正代码参考如下:
import torch import numpy as np from torch import nn import torch.nn.functional as F from torchsummary import summary class exampleNet(nn.Module): def __init__(self): super(exampleNet, self).__init__() # 注册为缓冲区,不参与梯度更新,会自动跟随模型切换设备、精度 w = torch.ones(1, 1, 2, 2) self.register_buffer('w', w) def forward(self, x): x = F.conv2d(x, self.w) return x model = exampleNet() summary(model, (1,10,10))
- 如果你有特殊需求一定要在forward中动态生成权重,生成后手动对齐输入的设备和精度即可,在
w = self.filters()后添加以下代码:
w = w.to(x.device, x.dtype)
内容的提问来源于stack exchange,提问作者stish
相关产品推荐
相关产品推荐

