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

如何让PyTorch CNN同时支持单张与多张MNIST图片输入?

解决CNN模型同时支持单张和批量图片输入的问题

我在为MNIST数据集构建CNN模型,训练过程正常,但输入单张图片时出现形状匹配错误;如果修改全连接层的in_features适配单张图片,训练循环又无法正常运行。想请教如何让模型同时支持单张和多张图片输入?

模型代码

import torch.nn as nn

class Net(nn.Module):
    def __init__(self, input_channels: int, output_features: int):
        super().__init__()

        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels=input_channels,
                      out_channels=16,
                      kernel_size=(3, 3), # 卷积核的高和宽
                      stride=1, # 卷积核移动的步长
                      padding=1), # 边缘填充,保证输出特征图尺寸和输入一致

            nn.ReLU(),

            nn.Conv2d(in_channels=16,
                      out_channels=16,
                      kernel_size=(3, 3),
                      stride=1,
                      padding=1),

            nn.ReLU(),

            nn.MaxPool2d(kernel_size=(2, 2)) # 池化层缩小特征图尺寸
        )

        self.conv2 = nn.Sequential(
            nn.Conv2d(in_channels=16,
                      out_channels=16,
                      kernel_size=(3, 3),
                      stride=1,
                      padding=1),

            nn.ReLU(),

            nn.Conv2d(in_channels=16,
                      out_channels=16,
                      kernel_size=(3, 3),
                      stride=1,
                      padding=1),

            nn.ReLU(),

            nn.MaxPool2d(kernel_size=(2, 2))
        )

        self.fc1 = nn.Sequential(
            nn.Flatten(),
            nn.Linear(in_features=16 * 7 * 7, out_features=output_features)
        )

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.fc1(x)

        return x

model = Net(input_channels=1 ,output_features=10)

训练代码

from tqdm.auto import tqdm
import time

start = time.time()

model.to(device)

epochs = 3

for epoch in range(epochs):
    print(f"Epoch: {epoch}\n---------")
    train_loss = 0

    # 训练模式
    model.train()

    for batch, (X, y) in enumerate(train_dataloader):
        X, y = X.to(device), y.to(device)

        y_logits = model(X)

        loss = loss_fn(y_logits, y)

        train_loss += loss.item()

        optimizer.zero_grad()

        loss.backward()

        optimizer.step()

    train_loss /= len(train_dataloader)

    print(f"Cross Entropy Train Loss: {train_loss: .5f}")

    # 测试模式
    test_loss = 0

    model.eval()

    with torch.inference_mode():
        for batch, (X, y) in enumerate(test_dataloader):
            y_logits = model(X)

            loss = loss_fn(y_logits, y)

            test_loss += loss.item()

        test_loss /= len(test_dataloader)

    print(f"Cross Entropy Test Loss: {test_loss: .5f}")

end = time.time()

print(f"Train Time on {device.upper()}, {round(end-start, 5)} seconds")

报错信息

输入单张图片的代码:

image, label = train[0]
model(image)

报错内容:

/usr/local/lib/python3.10/dist-packages/torch/nn/modules/linear.py in forward(self, input)
    114 
    115     def forward(self, input: Tensor) -> Tensor:
--> 116         return F.linear(input, self.weight, self.bias)
    117 
    118     def extra_repr(self) -> str:

RuntimeError: mat1 and mat2 shapes cannot be multiplied (16x49 and 784x10)

问题原因

PyTorch中CNN层(如Conv2d)要求输入形状为**[batch_size, channels, height, width]**:

  • 训练时,train_dataloader输出的批量数据符合这个格式(例如[64,1,28,28],64是批量大小),经过卷积池化后得到[64,16,7,7],展平后是[64, 16*7*7=784],和全连接层in_features=784匹配。
  • 但单张图片train[0]的形状是[1,28,28](缺少batch维度),输入模型后,卷积层会把第一个维度当作batch维度,通道维度被错误识别,最终展平后的特征形状为[16,49],和全连接层的784输入特征数不匹配,导致报错。

解决方法

方法1:输入单张图片时手动添加batch维度

调用模型前,用unsqueeze(0)给图片增加batch维度,让输入形状从[1,28,28]变为[1,1,28,28],和训练时的输入格式一致:

image, label = train[0]
# 添加batch维度
model(image.unsqueeze(0))

方法2:修改模型自动适配输入维度

在模型的forward方法中,自动检测输入维度,若为3维(无batch)则添加batch维度,这样无需手动修改输入:

def forward(self, x):
    # 如果输入是3维(channels, height, width),添加batch维度
    if x.dim() == 3:
        x = x.unsqueeze(0)
    x = self.conv1(x)
    x = self.conv2(x)
    x = self.fc1(x)
    return x

两种方法都能让模型同时支持单张和批量图片输入,推荐方法2,使用更便捷。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 09:16:01