RuntimeError通道不匹配:如何修改PyTorch模型适配3通道输入?
解决方案
报错原因
模型第一个卷积层conv1定义的输入通道数为1,但你输入的是3通道RGB图像,导致维度不匹配触发RuntimeError。同时原模型的全连接层是针对小尺寸输入(如MNIST的28×28)设计的,需要适配224×224的输入尺寸。
修改步骤
调整第一个卷积层的输入通道
将conv1的输入通道从1改为3,匹配3通道输入:self.conv1 = nn.Conv2d(3, 20, 5, 1)重新计算全连接层的输入维度
针对224×224的输入,逐层推导特征图尺寸:- 经过
conv1(核大小5,步长1,无padding):224 - 5 + 1 = 220,再经过max_pool2d(2,2)后尺寸变为220//2 = 110 - 经过
conv2(核大小5,步长1):110 - 5 + 1 = 106,再经过max_pool2d(2,2)后尺寸变为106//2 = 53 - 最终特征图尺寸为
53×53,通道数50,因此全连接层fc1的输入维度应为53*53*50
- 经过
修改全连接层与视图转换
更新fc1的输入特征数,并同步调整view的维度:self.fc1 = nn.Linear(53*53*50, 500) # forward中对应修改 x = x.view(-1, 53*53*50)
修改后的完整代码
import torch import torch.nn as nn import torch.nn.functional as F class Net(nn.Module): def __init__(self): super(Net, self).__init__() # 适配3通道输入 self.conv1 = nn.Conv2d(3, 20, 5, 1) self.conv2 = nn.Conv2d(20, 50, 5, 1) # 适配224×224输入的全连接层维度 self.fc1 = nn.Linear(53*53*50, 500) self.fc2 = nn.Linear(500, 10) def forward(self, x): x = F.relu(self.conv1(x)) x = F.max_pool2d(x, 2, 2) x = F.relu(self.conv2(x)) x = F.max_pool2d(x, 2, 2) x = x.view(-1, 53*53*50) x = F.relu(self.fc1(x)) x = self.fc2(x) return F.log_softmax(x, dim=1)
可选方案(不推荐)
如果想保留原模型的全连接层结构,可在conv1前添加1×1卷积层将3通道转为1通道,但会丢失RGB通道信息:
# __init__中新增 self.conv0 = nn.Conv2d(3, 1, 1, 1) # forward中先执行 x = self.conv0(x)
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

