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

PyTorch CNN训练批次尺寸不匹配错误的解决方法求助

PyTorch CNN训练批次尺寸不匹配错误的解决方法求助

问题描述

大家好,我现在在用PyTorch训练一个针对场景分类的CNN模型,优化器用的是SGD,但训练循环里一直弹出「Expected input batchsize to match target batchsize」的错误。我尝试过调整数据加载逻辑和模型输入维度,但问题还是没解决,想请教下大家该怎么正确处理训练循环里的批次尺寸,彻底解决这个不匹配的问题?

我的数据集情况:有一个dataset文件夹,里面包含4个子文件夹(forest、glacies、mountain、sea),每个子文件夹里大概有25000张对应场景的jpg图片。

我的代码

import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
import os
from PIL import Image
from sklearn.model_selection import train_test_split
import numpy as np
import matplotlib.pyplot as plt
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F

class ConvNet(nn.Module):
    def __init__(self, num_classes=4):
        super(ConvNet, self).__init__()
        # Convolutional layers
        self.conv1 = nn.Conv2d(in_channels=3, out_channels=4, kernel_size=3, stride=1, padding=1)
        self.conv2 = nn.Conv2d(in_channels=4, out_channels=8, kernel_size=3, stride=1, padding=1)
        self.conv3 = nn.Conv2d(in_channels=8, out_channels=16, kernel_size=3, stride=1, padding=1)
        # Max-pooling layers
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
        # Fully connected (linear) layer
        self.fc = nn.Linear(16 * 64 * 64, num_classes)  # Adjust the input size based on your image dimensions

    def forward(self, X):
        # Convolutional layers with ReLU activations and max-pooling
        X = F.relu(self.conv1(X))
        X = self.pool(X)
        X = F.relu(self.conv2(X))
        X = self.pool(X)
        X = F.relu(self.conv3(X))
        X = self.pool(X)
        # Flatten the output for the fully connected layer
        X = X.view(-1, 16 * 64 * 64)  # Adjust the size based on your image dimensions
        # Fully connected layer
        X = self.fc(X)
        return X

class SceneDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.image_list, self.labels = self.load_dataset()
        # Create a mapping from class names to indices
        self.class_to_index = {class_name: idx for idx, class_name in enumerate(set(self.labels))}

    def load_dataset(self):
        image_list = []
        labels = []
        for class_name in os.listdir(self.root_dir):
            class_path = os.path.join(self.root_dir, class_name)
            if os.path.isdir(class_path):
                label = class_name
                for filename in os.listdir(class_path):
                    if filename.endswith(".jpg"):
                        image_list.append(os.path.join(class_path, filename))
                        labels.append(label)
        return image_list, labels

    def __len__(self):
        return len(self.image_list)

    def __getitem__(self, index):
        img_path = self.image_list[index]
        label = self.labels[index]
        image = Image.open(img_path).convert('RGB')
        if self.transform:
            image = self.transform(image)
        # Get the class index
        label_index = self.class_to_index[label]
        # Convert label to tensor
        label_tensor = torch.tensor(label_index, dtype=torch.long)
        return image, label_tensor

def get_dataloaders(root, train_batchsize, test_batchsize):
    transform = transforms.Compose([
        transforms.Resize((256, 256)),
        transforms.ToTensor(),
    ])
    dataset = SceneDataset(root, transform=transform)
    # Split the dataset into train, validation, and test sets
    train_size = int(0.7 * len(dataset))
    val_size = int(0.1 * len(dataset))
    test_size = len(dataset) - train_size - val_size
    train_dataset, val_dataset, test_dataset = torch.utils.data.random_split(
        dataset, [train_size, val_size, test_size])
    # Create data loaders
    train_dataloader = DataLoader(train_dataset, batch_size=train_batchsize, shuffle=True)
    val_dataloader = DataLoader(val_dataset, batch_size=test_batchsize, shuffle=False)
    test_dataloader = DataLoader(test_dataset, batch_size=test_batchsize, shuffle=False)
    return train_dataloader, val_dataloader, test_dataloader

# Example usage
root_directory = "data"
train_batchsize = 32
test_batchsize = 1
train_dataloader, val_dataloader, test_dataloader = get_dataloaders(root_directory, train_batchsize, test_batchsize)

# Helper for visualization
def img_show(image, label):
    plt.figure()
    plt.title(f'This is a {label}')
    im = np.moveaxis(np.array(image), [0,1,2], [2, 0, 1])
    plt.imshow(im)
    plt.show()

# Visualize first 4 samples
for count, (image, label) in enumerate(train_dataloader):
    img_show(image[0], label[0])
    if count == 3:
        break

max_epoch = 300
train_batch = 32
test_batch = 1
learning_rate = 0.01

# Create train, validation, and test dataset loaders
train_loader, val_loader, test_loader = get_dataloaders(root_directory, train_batch, train_batch)  # Use the same batch size for validation

# Initialize your network
model = ConvNet()

# Define your loss function
criterion = nn.CrossEntropyLoss()

# Initialize optimizer
optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate, weight_decay=5e-04)

# Placeholder for best validation accuracy
best_val_accuracy = 0.0

# Placeholder for the best model state
best_model_state = None

# Placeholder for training and validation statistics
train_losses, val_losses = [], []
train_accuracies, val_accuracies = [], []

# Start training
for epoch in range(max_epoch):
    model = model.train()
    total_train_loss = 0.0
    correct_train = 0
    total_train = 0
    for images, labels in train_loader:
        optimizer.zero_grad()
        # Forward pass
        outputs = model(images)
        # Ensure labels have the correct shape
        if labels.size(0) != outputs.size(0):
            labels = labels[:outputs.size(0)]
        loss = criterion(outputs, labels.squeeze().long())  # Adjusted for label size
        # Backward pass and optimization
        loss.backward()
        optimizer.step()
        total_train_loss += loss.item()
        _, predicted = torch.max(outputs.data, 1)
        print(f"Predicted shape: {predicted.shape}, Labels shape: {labels[:predicted.size(0)].squeeze().shape}")
        total_train += labels.size(0)
        batch_size = min(labels.size(0), predicted.size(0))
        correct_train += (predicted[:batch_size] == labels[:batch_size].squeeze()).sum().item()

    # Calculate training accuracy and loss
    train_accuracy = correct_train / total_train
    train_losses.append(total_train_loss / len(train_loader))
    train_accuracies.append(train_accuracy)

    # Validation
    model = model.eval()
    total_val_loss = 0.0
    correct_val = 0
    total_val = 0
    with torch.no_grad():
        for images, labels in val_loader:
            outputs = model(images)
            loss = criterion(outputs, labels.squeeze().long()) # Convert labels to long tensor
            total_val_loss += loss.item()
            _, predicted = torch.max(outputs.data, 1)
            total_train += labels.size(0)
            correct_train += (predicted == labels[:predicted.size(0)].squeeze()).sum().item()

    # Calculate validation accuracy and loss
    val_accuracy = correct_val / total_val
    val_losses.append(total_val_loss / len(val_loader))
    val_accuracies.append(val_accuracy)

    # Save the best model based on validation accuracy
    if val_accuracy > best_val_accuracy:
        best_val_accuracy = val_accuracy
        best_model_state = model.state_dict()

# Save the best model state to a file
best_model_path = "best_cnn_sgd.pth"
torch.save(best_model_state, best_model_path)

# Plot losses vs epoch
plt.figure(figsize=(10, 5))
plt.plot(train_losses, label='Training Loss')
plt.plot(val_losses, label='Validation Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Training and Validation Loss vs. Epoch')
plt.legend()
plt.show()

# Plot accuracies vs epoch
plt.figure(figsize=(10, 5))
plt.plot(train_accuracies, label='Training Accuracy')
plt.plot(val_accuracies, label='Validation Accuracy')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.title('Training and Validation Accuracy vs. Epoch')
plt.legend()
plt.show()

我已经注意到训练循环里的批次尺寸有问题,导致预测批次和目标批次不匹配,也尝试过调整数据加载和模型输入维度,但问题还是存在,恳请大家给我一些指导。


解答

Hey,我帮你排查了代码里的几个关键问题,这个批次不匹配的错误主要是由以下几个原因导致的,咱们逐个修正:

1. 模型特征图尺寸计算错误(核心问题)

你的输入图片被Resize到了(256,256),然后经过了三次MaxPool2d(kernel_size=2, stride=2)(每次池化都会把尺寸减半):

  • 第一次池化后:256 → 128
  • 第二次池化后:128 → 64
  • 第三次池化后:64 → 32

但你在模型的全连接层和view操作里写的是16 * 64 * 64,这会导致forward过程中张量形状不匹配,进而输出的outputs批次维度异常,和labels的批次尺寸对不上。

修正方法:把模型里的尺寸改成16 * 32 * 32:

class ConvNet(nn.Module):
    def __init__(self, num_classes=4):
        super(ConvNet, self).__init__()
        self.conv1 = nn.Conv2d(in_channels=3, out_channels=4, kernel_size=3, stride=1, padding=1)
        self.conv2 = nn.Conv2d(in_channels=4, out_channels=8, kernel_size=3, stride=1, padding=1)
        self.conv3 = nn.Conv2d(in_channels=8, out_channels=16, kernel_size=3, stride=1, padding=1)
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
        # 修正全连接层输入尺寸
        self.fc = nn.Linear(16 * 32 * 32, num_classes)  

    def forward(self, X):
        X = F.relu(self.conv1(X))
        X = self.pool(X)
        X = F.relu(self.conv2(X))
        X = self.pool(X)
        X = F.relu(self.conv3(X))
        X = self.pool(X)
        # 修正flatten的尺寸
        X = X.view(-1, 16 * 32 * 32)  
        X = self.fc(X)
        return X

2. 验证循环里的变量复制错误

在验证阶段的代码里,你错误地使用了total_train和correct_train来累计验证集的统计数据,而不是total_val和correct_val,这会导致val_accuracy计算出错(分母可能为0),也会干扰训练集的统计结果。

修正方法:把验证循环里的变量替换成正确的:

# Validation
model = model.eval()
total_val_loss = 0.0
correct_val = 0
total_val = 0
with torch.no_grad():
    for images, labels in val_loader:
        outputs = model(images)
        # 不需要额外处理标签,Dataset已经返回正确的long类型张量
        loss = criterion(outputs, labels)
        total_val_loss += loss.item()
        _, predicted = torch.max(outputs.data, 1)
        total_val += labels.size(0)
        correct_val += (predicted == labels).sum().item()

# 计算验证准确率和损失
val_accuracy = correct_val / total_val
val_losses.append(total_val_loss / len(val_loader))
val_accuracies.append(val_accuracy)

3. 冗余的标签处理导致维度异常

你的SceneDataset已经在__getitem__里返回了torch.long类型的标签张量,DataLoader会自动把这些单样本标签堆叠成[batch_size]形状的张量。但你在训练和验证循环里额外做了labels.squeeze().long(),当批次尺寸为1时,squeeze()会把标签变成0维张量,和outputs的[1,4]形状不匹配,引发批次尺寸错误。

修正方法:直接使用原始的labels,去掉冗余的处理:

# 训练循环里的损失计算修改为:
loss = criterion(outputs, labels)

4. 移除不必要的批次尺寸判断代码

你之前加的if labels.size(0) != outputs.size(0): labels = labels[:outputs.size(0)]这类代码是冗余的,只要模型输入输出维度正确,DataLoader返回的images和labels的批次尺寸肯定一致,直接删掉这些代码即可。


把这些修改全部应用后,你的训练循环应该就能正常运行,批次尺寸不匹配的错误也会消失。另外建议你可以在模型的forward函数里打印一下各层的张量形状,方便后续排查类似的维度问题。

备注:内容来源于stack exchange,提问作者Furkan Şenoğlu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.21 07:08:06