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

自定义图像分类模型训练报错:Conv2d通道维度不匹配求助

图像分类模型通道不匹配错误排查与修复

错误信息

给定groups=1,权重尺寸为[16, 32, 3, 3],期望输入[42, 19, 224, 224]有32个通道,但实际得到19个通道

问题根源

你的模型存在两处硬编码导致的维度不匹配问题:

  1. CNN模块卷积通道硬编码:
    • 拼接后的combined_features通道数是input_channels + hidden_units(当前参数下为3+16=19),但self.cnn中第二个卷积层的输入通道写死为19,而它的前一层输出是hidden_units=16,直接导致输入输出通道不匹配。
    • 第一个卷积层的输入通道也写死为19,虽然当前参数下刚好匹配,但后续修改输入通道或隐藏单元数会再次出错。
  2. 全连接层输入维度硬编码:
    • self.fc的输入维度写死为19*2*56*56,完全依赖固定输入尺寸,缺乏灵活性且极易和实际输出维度不匹配。
  3. 冗余代码:forward函数里有一行单独的global_features,属于无效代码,可直接删除。

修复方案

步骤1:修正CNN模块的卷积通道

将self.cnn中的硬编码通道替换为动态参数,确保前后层通道匹配:

self.cnn = nn.Sequential(
    # 第一个卷积层输入通道为拼接后的总通道数:input_channels + hidden_units
    nn.Conv2d(input_channels + hidden_units, hidden_units, kernel_size=3, padding=1),
    nn.ReLU(inplace=True),
    nn.MaxPool2d(kernel_size=2, stride=2),
    # 第二个卷积层输入通道为前一层的输出通道hidden_units
    nn.Conv2d(hidden_units, hidden_units * 2, kernel_size=3, padding=1),
    nn.ReLU(inplace=True),
    nn.MaxPool2d(kernel_size=2, stride=2)
)

步骤2:修正全连接层输入维度

使用自适应池化让特征图尺寸自适应,彻底避免硬编码尺寸:

# 替换原CNN模块和全连接层初始化代码
self.cnn = nn.Sequential(
    nn.Conv2d(input_channels + hidden_units, hidden_units, kernel_size=3, padding=1),
    nn.ReLU(inplace=True),
    nn.MaxPool2d(kernel_size=2, stride=2),
    nn.Conv2d(hidden_units, hidden_units * 2, kernel_size=3, padding=1),
    nn.ReLU(inplace=True),
    nn.MaxPool2d(kernel_size=2, stride=2),
    # 添加自适应平均池化,将任意尺寸特征图转为1x1
    nn.AdaptiveAvgPool2d((1, 1))
)
# 全连接层输入维度改为hidden_units*2(自适应池化后的特征数)
self.fc = nn.Linear(hidden_units * 2, output_classes)

同时修改forward中flatten部分的代码(无需改动,自适应池化后直接展开即可):

cnn_output = cnn_output.view(cnn_output.size(0), -1)

步骤3:删除冗余代码

移除forward函数中单独的global_features空行。

完整修正后的代码

import torch
import torch.nn as nn
import torch.nn.functional as F

class CustomModel(nn.Module):
    def __init__(self, input_channels, hidden_units, output_classes):
        super(CustomModel, self).__init__()
        
        # Global feature extraction layers
        self.global_pool = nn.AdaptiveAvgPool2d((1, 1))  # Global average pooling
        
        # Local feature extraction layers
        self.local_conv = nn.Conv2d(input_channels, hidden_units, kernel_size=3, padding=1)
        
        # Convolutional neural network
        self.cnn = nn.Sequential(
            nn.Conv2d(input_channels + hidden_units, hidden_units, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=2, stride=2),
            nn.Conv2d(hidden_units, hidden_units * 2, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=2, stride=2),
            nn.AdaptiveAvgPool2d((1, 1))
        )
        
        # Fully connected layer
        self.fc = nn.Linear(hidden_units * 2, output_classes)

    def forward(self, x):
        # Global feature extraction
        global_features = self.global_pool(x)
        global_features = global_features.view(global_features.size(0), -1)
        global_features = global_features.unsqueeze(-1).unsqueeze(-1)  # Expand dimensions to match local features
        global_features = global_features.expand(-1, -1, x.size(2), x.size(3))  # Expand to match spatial dimensions
        
        # Local feature extraction
        local_features = self.local_conv(x)
        local_features = F.relu(local_features)
        
        # Concatenate global and local features
        combined_features = torch.cat((global_features, local_features), dim=1)
        
        # CNN processing
        cnn_output = self.cnn(combined_features)
        
        # Flatten for fully connected layer
        cnn_output = cnn_output.view(cnn_output.size(0), -1)
        
        # Fully connected layer
        output = self.fc(cnn_output)
        return output
# 模型初始化代码
input_channels = 3
hidden_units = 16
output_classes = 75  # 对应原len(class_names)的值
custom_model = CustomModel(input_channels, hidden_units, output_classes)

内容的提问来源于stack exchange,提问作者Stand Still

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 08:06:05