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

添加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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 13:35:44