添加Conv3后CIFAR10模型出现矩阵乘法形状不匹配问题求助
问题分析与解决方案
错误本质:特征图维度与全连接层不匹配
你碰到的RuntimeError: mat1 and mat2 shapes cannot be multiplied (4x32 and 400x120),核心原因是卷积层输出的特征图展平后,维度和全连接层的输入维度不匹配:
mat1是展平后的特征张量,形状(batch_size, 32)(这里batch_size为4)mat2是全连接层fc1的权重矩阵,要求输入维度必须是400,但实际传入的是32,矩阵乘法无法执行。
关键:计算卷积层输出的特征图尺寸
卷积层输出尺寸(高度/宽度方向)的计算公式:
# H_out = (H_in - kernel_size + 2*padding) // stride + 1 # W_out 计算逻辑和H_out完全一致
结合CIFAR10的3x32x32输入,以及你新增的conv3,我们一步步推导维度变化:
假设你的基础网络(新增conv3前)是常见结构:
self.conv1 = nn.Conv2d(3, 16, 3, padding=1) # 输出:16x32x32 self.pool = nn.MaxPool2d(2, 2) # 池化后:16x16x16 self.conv2 = nn.Conv2d(16, 16, 3, padding=1) # 输出:16x16x16 self.pool(F.relu(self.conv2(x))) # 池化后:16x8x8
现在新增self.conv3 = nn.Conv2d(16, 32, 5)(默认padding=0、stride=1),代入公式计算:
H_out = (8 - 5 + 0) // 1 + 1 = 4 W_out = (8 -5 +0) //1 +1 =4
所以conv3输出的特征图形状是32x4x4,展平后一维维度为32*4*4=512。
解决方案:修正全连接层维度并定义fc1_5
步骤1:快速定位展平维度
可以在forward函数中打印每一层的输出形状,直接确认展平后的维度:
def forward(self, x): x = self.pool(F.relu(self.conv1(x))) print("conv1+pool shape:", x.shape) x = self.pool(F.relu(self.conv2(x))) print("conv2+pool shape:", x.shape) x = F.relu(self.conv3(x)) print("conv3 shape:", x.shape) x = x.view(-1, 32*4*4) # 这里的数值需要和conv3输出的通道*高*宽一致 print("flatten shape:", x.shape) # 后续全连接层逻辑 return x
运行代码后,flatten shape的第二个数值就是全连接层的输入维度。
步骤2:修改fc1的定义
比如上面推导的展平维度是512,那么fc1应该定义为:
self.fc1 = nn.Linear(512, 120) # 输入512,输出120(保持你原来的输出维度)
定义self.fc1_5
fc1_5是在fc1和最终输出层之间新增的隐藏层,作用是进一步提取高级特征,定义逻辑和其他全连接层一致:
- 输入维度等于
fc1的输出维度(比如120) - 输出维度可根据需求自行设定(比如64)
示例代码:
self.fc1 = nn.Linear(512, 120) self.fc1_5 = nn.Linear(120, 64) # 输入120,输出64 self.fc2 = nn.Linear(64, 10) # 最终输出10类(CIFAR10的类别数)
对应的forward函数需要新增fc1_5的前向传播:
def forward(self, x): # 卷积层部分 x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = F.relu(self.conv3(x)) x = x.view(-1, 32*4*4) # 全连接层部分 x = F.relu(self.fc1(x)) x = F.relu(self.fc1_5(x)) x = self.fc2(x) return x
完整可运行示例
import torch import torch.nn as nn import torch.nn.functional as F class CIFAR10Net(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(3, 16, 3, padding=1) self.pool = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(16, 16, 3, padding=1) self.conv3 = nn.Conv2d(16, 32, 5) self.fc1 = nn.Linear(32*4*4, 120) self.fc1_5 = nn.Linear(120, 64) self.fc2 = nn.Linear(64, 10) def forward(self, x): x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = F.relu(self.conv3(x)) x = x.view(-1, 32*4*4) x = F.relu(self.fc1(x)) x = F.relu(self.fc1_5(x)) x = self.fc2(x) return x # 测试网络 net = CIFAR10Net() test_input = torch.randn(4, 3, 32, 32) test_output = net(test_input) print("输出形状:", test_output.shape) # 应为 torch.Size([4, 10])
内容的提问来源于stack exchange,提问作者Rebekah Lee
相关产品推荐
相关产品推荐

