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

如何正确定义可适配任意输入尺寸的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 19:24:05