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

带残差连接的LeNet-5架构搭建报错:输入通道不匹配排查

问题分析与修复方案

你的代码存在几个核心问题,直接导致了维度不匹配的报错,以下是具体问题和修正后的完整代码:

核心问题点

  • ResidualBlock未实现forward方法:自定义Module必须实现前向传播逻辑,否则无法完成张量的正向计算。
  • MyModel的forward错误访问子模块内部属性:你直接调用self.conv、self.activation,但这些是ResidualBlock类的内部属性,应该调用实例化的残差块对象(如self.conv1(x))。
  • 卷积转全连接时未展平特征:卷积层输出是4维张量(batch, channels, h, w),必须展平为2维(batch, features)才能输入全连接层。
  • Softmax未指定维度:默认行为可能导致错误的归一化维度,需显式指定dim=1在类别维度计算概率。

修正后的代码

1. 修复ResidualBlock类

import torch

class ResidualBlock(torch.nn.Module): 

    def __init__(  
        self,
        input_c,  # 输入通道数
        output_c,  # 输出通道数
        kernel_size,  # 卷积核尺寸
        activation=torch.nn.ReLU,  # 激活函数类
    ):
        super().__init__()
        # 用于引入非线性的激活函数
        self.activation = activation()

        # 选择合适的填充,使图像的高和宽尺寸保持不变
        # (假设卷积核尺寸始终为奇数)
        padding_size = (kernel_size - 1) // 2

        # 卷积操作
        self.conv = torch.nn.Conv2d(
            in_channels=input_c,
            out_channels=output_c,
            kernel_size=kernel_size,
            padding=padding_size,
            padding_mode="zeros",
        )

        # 若输入与输出通道数不同,则使用1x1卷积进行通道调整
        if input_c != output_c:
            self.correct_channels = torch.nn.Conv2d(
                in_channels=input_c,
                out_channels=output_c,
                kernel_size=1,
            )
        else:
            self.correct_channels = torch.nn.Identity()

    # 新增forward方法,实现残差连接核心逻辑
    def forward(self, x):
        conv_out = self.conv(x)
        residual = self.correct_channels(x)
        return self.activation(conv_out + residual)

2. 修复MyModel类

class MyModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        
        self.conv1 = ResidualBlock(input_c=1, output_c=6, kernel_size=5)
        self.maxpool1 = torch.nn.MaxPool2d(kernel_size=2)
        self.conv2 = ResidualBlock(input_c=6, output_c=16, kernel_size=5)
        self.maxpool2 = torch.nn.MaxPool2d(kernel_size=2)
        self.conv3 = ResidualBlock(input_c=16, output_c=120, kernel_size=5)

        self.lin1 = torch.nn.Linear(in_features=120, out_features=84)
        self.relu1 = torch.nn.ReLU()
        self.lin2 = torch.nn.Linear(in_features=84, out_features=10)
        # 指定dim=1,在类别维度做概率归一化
        self.softmax = torch.nn.Softmax(dim=1)

    def forward(self, x):
        # 调用实例化的残差块,而非直接访问子模块内部属性
        x = self.conv1(x)
        x = self.maxpool1(x)
        x = self.conv2(x)
        x = self.maxpool2(x)
        x = self.conv3(x)

        # 展平特征:从(batch, 120, 1, 1)转为(batch, 120)
        x = torch.flatten(x, start_dim=1)
        
        x = self.lin1(x)
        x = self.relu1(x)
        x = self.lin2(x)
        x = self.softmax(x)

        return x

# 测试模型可行性
model = MyModel()
# 输入MNIST格式的张量:(batch_size, channels, h, w)
test_input = torch.randn(1, 1, 28, 28)
output = model(test_input)
print(output.shape)  # 应输出torch.Size([1, 10])

验证说明

修正后,输入MNIST格式的张量(1通道、28x28)可以正常完成前向传播,输出10类的概率分布,维度不匹配的报错完全解决。

内容的提问来源于stack exchange,提问作者2Razor007

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 15:27:07