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

PyTorch中LeNet-5适配CIFAR10的参数调整问题(固定batch size=16)

问题分析与解决

错误1:特征图维度计算错误引发的RuntimeError

RuntimeError: shape '[-1, 2048]' is invalid for input of size 256 的核心原因是错误估算了卷积+池化后的特征图尺寸,导致全连接层的输入维度与实际张量总元素数不匹配。

特征图尺寸逐阶计算(基于你的卷积参数)

输入张量形状:(16, 3, 32, 32)(batch_size=16,3通道,32×32像素)

  • Conv1 + 池化:
    Conv1参数为 nn.Conv2d(3, 6, 5, stride=1, padding=1),用卷积输出尺寸公式:(输入尺寸 - 卷积核尺寸 + 2×padding) / stride + 1
    计算得:(32 - 5 + 2×1)/1 + 1 = 30 → 输出形状为 (16, 6, 30, 30)
    经过MaxPool2d(2,2)(步长2)后,尺寸减半 → 输出形状:(16, 6, 15, 15)
  • Conv2 + 池化:
    Conv2参数为 nn.Conv2d(6, 16, 5)(默认stride=1,padding=0),计算得:(15 -5)/1 +1 =11 → 输出形状:(16,16,11,11)
    经过MaxPool2d(2,2)后,尺寸计算为 (11-2)//2 +1 =5 → 输出形状:(16,16,5,5)
  • Flatten后总特征数:每个样本的特征数为 16×5×5=400,全连接层fc1的输入维度必须对应这个数值。

错误2:batch_size不匹配的ValueError

ValueError: Expected input batch_size (50) to match target batch_size (16) 是因为调整后的全连接层输入维度仍错误,x.view(-1, N) 中的-1会被自动计算为总元素数/N,当N与实际特征数不符时,这个计算结果就会和设定的batch_size=16冲突。

修正后的完整模型代码

import torch.nn as nn

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5, stride=1, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        # 修正全连接层输入维度:16*5*5=400
        self.fc1 = nn.Linear(16*5*5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)
        self.relu = nn.ReLU()

    def forward(self, x):
        x = self.pool(self.relu(self.conv1(x)))
        x = self.pool(self.relu(self.conv2(x)))
        # 对应全连接层的输入维度,保证batch_size=16不变
        x = x.view(-1, 16*5*5)  
        x = self.relu(self.fc1(x))
        x = self.relu(self.fc2(x))
        x = self.fc3(x)
        return x

net = Net()

维度验证方法

可以在forward函数中添加打印语句,确认每一步的张量形状是否符合预期:

def forward(self, x):
    print("初始输入:", x.shape)
    x = self.pool(self.relu(self.conv1(x)))
    print("Conv1+Pool后:", x.shape)
    x = self.pool(self.relu(self.conv2(x)))
    print("Conv2+Pool后:", x.shape)
    x = x.view(-1, 16*5*5)  
    print("Flatten后:", x.shape)
    x = self.relu(self.fc1(x))
    x = self.relu(self.fc2(x))
    x = self.fc3(x)
    return x

运行后会输出:

初始输入: torch.Size([16, 3, 32, 32])
Conv1+Pool后: torch.Size([16, 6, 15, 15])
Conv2+Pool后: torch.Size([16, 16, 5, 5])
Flatten后: torch.Size([16, 400])

这样就能确保batch_size始终为16,各层维度完全匹配。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 15:00:17