如何正确定义可适配任意输入尺寸的PyTorch CNN模型
问题原因
你遇到的RuntimeError: mat1 dim 1 must match mat2 dim 0错误本质是全连接层输入维度不匹配:原LeNet实现中linear1的输入维度被写死为120,仅当输入为(1,32,32)的灰度图时,经过3次卷积、2次池化后输出的特征图尺寸为(120,1,1),拉平后刚好是120维匹配全连接层输入。当输入改为6464图像时,conv3输出的特征图尺寸大于11,拉平后的维度远大于120,因此触发矩阵相乘的维度不匹配错误。
PyTorch基于动态图机制,不需要像TensorFlow静态图那样提前通过Input层固定输入尺寸,只要保证前向传播每一步运算的维度匹配即可,以下是两种成熟的解决方案:
解决方案
方案1:使用自适应池化层(推荐,支持任意输入尺寸)
自适应池化层可以在不指定输入尺寸的前提下,固定输出特征图的尺寸,不需要提前计算卷积后的维度,也不需要手动调整全连接层参数,适配任意输入图像尺寸:
import torch import torch.nn as nn class AdaptiveLeNet(nn.Module): # 可根据需求调整输入通道数、分类数,默认适配3通道RGB猫狗二分类 def __init__(self, in_channels=3, num_classes=2): super().__init__() self.relu = nn.ReLU() self.pool = nn.AvgPool2d(kernel_size=2, stride=2) self.conv1 = nn.Conv2d(in_channels=in_channels, out_channels=6, kernel_size=5, stride=1, padding=0) self.conv2 = nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5, stride=1, padding=0) self.conv3 = nn.Conv2d(in_channels=16, out_channels=120, kernel_size=5, stride=1, padding=0) # 新增自适应平均池化,固定输出特征图尺寸为1*1 self.adaptive_pool = nn.AdaptiveAvgPool2d(output_size=(1,1)) self.linear1 = nn.Linear(120, 84) self.linear2 = nn.Linear(84, num_classes) def forward(self, x): x = self.relu(self.conv1(x)) x = self.pool(x) x = self.relu(self.conv2(x)) x = self.pool(x) x = self.relu(self.conv3(x)) # 无论输入尺寸多大,输出都固定为(batch_size, 120, 1, 1) x = self.adaptive_pool(x) x = x.flatten(start_dim=1) x = self.relu(self.linear1(x)) x = self.linear2(x) return x
方案2:初始化时动态计算全连接层维度(固定输入尺寸场景适用)
如果训练/推理的输入尺寸是固定的,不想引入自适应池化层,也可以在模型初始化阶段传入输入尺寸参数,通过dummy张量自动计算全连接层的输入维度,避免手动计算误差:
import torch import torch.nn as nn class DynamicLenet(nn.Module): # 初始化时传入输入尺寸即可自动适配,比如64*64灰度图传入input_size=(1,64,64) def __init__(self, in_channels=3, num_classes=2, input_size=(3, 32, 32)): super().__init__() self.relu = nn.ReLU() self.pool = nn.AvgPool2d(kernel_size=2, stride=2) self.conv1 = nn.Conv2d(in_channels=in_channels, out_channels=6, kernel_size=5, stride=1, padding=0) self.conv2 = nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5, stride=1, padding=0) self.conv3 = nn.Conv2d(in_channels=16, out_channels=120, kernel_size=5, stride=1, padding=0) # 用伪输入自动计算卷积部分输出的维度 with torch.no_grad(): dummy_x = torch.randn(1, *input_size) conv_out = self.conv3(self.pool(self.conv2(self.pool(self.conv1(dummy_x))))) flatten_dim = conv_out.flatten(start_dim=1).shape[1] self.linear1 = nn.Linear(flatten_dim, 84) self.linear2 = nn.Linear(84, num_classes) def forward(self, x): x = self.relu(self.conv1(x)) x = self.pool(x) x = self.relu(self.conv2(x)) x = self.pool(x) x = self.relu(self.conv3(x)) x = x.flatten(start_dim=1) x = self.relu(self.linear1(x)) x = self.linear2(x) return x
内容的提问来源于stack exchange,提问作者Marcus
相关产品推荐
相关产品推荐

