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

PyTorch卫星图像分割中输入与目标batch_size不匹配问题求助

卫星图像分割代码中Batch Size不匹配问题排查与解决

问题描述

运行卫星图像7类分割代码时,出现如下错误:

ValueError: Expected input batch_size (3) to match target batch_size (9).

无论调整batch_size,目标batch_size始终是输入的3倍,无法定位问题根源。

问题根源

  1. 标签被错误应用图像归一化变换:代码中对标签和图像使用了相同的transform(包含Normalize),导致单通道的类别掩码被转换成3通道张量。后续labels.view(-1, 512, 512)会将[3,3,512,512]的张量重塑为[9,512,512],直接导致目标batch_size变为输入的3倍。
  2. 模型输出与标签尺寸不匹配:模型经过三次MaxPool2d(每次下采样2倍),输入512x512图像最终输出尺寸为64x64,但标签始终保持512x512,这会进一步引发维度不匹配问题。
  3. 冗余的Softmax层:CrossEntropyLoss内部已集成LogSoftmax计算,手动添加torch.softmax会导致损失计算逻辑错误。

修复步骤

1. 分离图像与标签的变换逻辑

图像保留归一化处理,标签仅转换为张量(无需归一化):

# 图像变换:包含归一化
image_transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
# 标签变换:仅转张量,保证单通道
label_transform = transforms.Compose([
    transforms.ToTensor(),
])

2. 修改Dataset类,分别应用变换

在__getitem__中对图像和标签使用各自的变换,并移除标签的冗余通道维度:

class SatelliteDataset(Dataset):
    def __init__(self, image_folder, label_folder, image_transform=None, label_transform=None):
        self.image_folder = image_folder
        self.label_folder = label_folder
        self.image_transform = image_transform
        self.label_transform = label_transform

        self.image_paths = sorted([os.path.join(image_folder, filename) for filename in os.listdir(image_folder)])
        self.label_paths = sorted([os.path.join(label_folder, filename) for filename in os.listdir(label_folder)])

    def __getitem__(self, idx):
        image_path = self.image_paths[idx]
        label_path = self.label_paths[idx]

        image = Image.open(image_path)
        label = Image.open(label_path)

        image = self.adjust_brightness(image, brightness_factor=1.8)

        # Resize images to 512x512
        image = image.resize((512, 512), Image.BILINEAR)
        label = label.resize((512, 512), Image.NEAREST)

        # 分别应用变换
        if self.image_transform:
            image = self.image_transform(image)
        if self.label_transform:
            label = self.label_transform(label)
        # 去除标签的通道维度,变为[512,512]
        label = label.squeeze(0)

        return image, label
    # ... 其余方法保持不变

3. 调整模型输出尺寸与标签匹配

将模型修改为带上采样跳连接的结构,让输出回到512x512,保证和标签尺寸一致:

class SegmentationCNN(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(SegmentationCNN, self).__init__()

        # 下采样编码器
        self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=3, stride=1, padding=1)
        self.relu1 = nn.ReLU()
        self.maxpool1 = nn.MaxPool2d(kernel_size=2, stride=2)  # 512→256

        self.conv2 = nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1)
        self.relu2 = nn.ReLU()
        self.maxpool2 = nn.MaxPool2d(kernel_size=2, stride=2)  # 256→128

        self.conv3 = nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1)
        self.relu3 = nn.ReLU()
        self.maxpool3 = nn.MaxPool2d(kernel_size=2, stride=2)  # 128→64

        self.conv4 = nn.Conv2d(256, 512, kernel_size=3, stride=1, padding=1)
        self.relu4 = nn.ReLU()

        # 上采样解码器
        self.upconv1 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)  # 64→128
        self.conv5 = nn.Conv2d(512, 256, kernel_size=3, stride=1, padding=1)  # 拼接跳连接特征
        self.relu5 = nn.ReLU()

        self.upconv2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)  # 128→256
        self.conv6 = nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1)
        self.relu6 = nn.ReLU()

        self.upconv3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)  # 256→512
        self.conv7 = nn.Conv2d(128, out_channels, kernel_size=3, stride=1, padding=1)

    def forward(self, x):
        # 编码阶段
        x1 = self.relu1(self.conv1(x))
        x = self.maxpool1(x1)

        x2 = self.relu2(self.conv2(x))
        x = self.maxpool2(x2)

        x3 = self.relu3(self.conv3(x))
        x = self.maxpool3(x3)

        x = self.relu4(self.conv4(x))

        # 解码阶段(带跳连接)
        x = self.upconv1(x)
        x = torch.cat([x, x3], dim=1)
        x = self.relu5(self.conv5(x))

        x = self.upconv2(x)
        x = torch.cat([x, x2], dim=1)
        x = self.relu6(self.conv6(x))

        x = self.upconv3(x)
        x = torch.cat([x, x1], dim=1)
        x = self.conv7(x)

        return x

4. 统一训练与验证函数的标签处理

在evaluate函数中添加和train函数一致的标签类型转换:

def evaluate(model, dataloader, loss_fn, device):
    model.eval()
    running_loss = 0.0

    with torch.no_grad():
        for images, labels in dataloader:
            images = images.to(device)
            labels = labels.to(device)

            outputs = model(images)
            labels = labels.long()

            loss = loss_fn(outputs, labels)
            running_loss += loss.item() * images.size(0)

    epoch_loss = running_loss / len(dataloader.dataset)
    return epoch_loss

5. 更新数据集与DataLoader初始化

# 创建数据集实例
dataset = SatelliteDataset(
    image_folder=image_folder,
    label_folder=label_folder,
    image_transform=image_transform,
    label_transform=label_transform
)

# 使用默认collate_fn即可
batch_size = 3
train_dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
val_dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False)

完整修复后代码

import torch
import torch.nn as nn
import torchvision.transforms as transforms
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import os
from PIL import ImageEnhance

class SatelliteDataset(Dataset):
    def __init__(self, image_folder, label_folder, image_transform=None, label_transform=None):
        self.image_folder = image_folder
        self.label_folder = label_folder
        self.image_transform = image_transform
        self.label_transform = label_transform

        self.image_paths = sorted([os.path.join(image_folder, filename) for filename in os.listdir(image_folder)])
        self.label_paths = sorted([os.path.join(label_folder, filename) for filename in os.listdir(label_folder)])

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

    def __getitem__(self, idx):
        image_path = self.image_paths[idx]
        label_path = self.label_paths[idx]

        image = Image.open(image_path)
        label = Image.open(label_path)

        image = self.adjust_brightness(image, brightness_factor=1.8)

        # Resize images to 512x512
        image = image.resize((512, 512), Image.BILINEAR)
        label = label.resize((512, 512), Image.NEAREST)

        # Apply transformations separately
        if self.image_transform:
            image = self.image_transform(image)
        if self.label_transform:
            label = self.label_transform(label)
        
        # Remove channel dimension from label (from [1,512,512] to [512,512])
        label = label.squeeze(0)

        return image, label

    def adjust_brightness(self, image, brightness_factor=1.0):
        enhancer = ImageEnhance.Brightness(image)
        enhanced_image = enhancer.enhance(brightness_factor)
        return enhanced_image


# Define paths to the folders containing satellite images and corresponding labels
image_folder = "/Users/.../train_data/images"
label_folder = "/Users/.../train_data/masks"

# Separate transforms for image and label
image_transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])

label_transform = transforms.Compose([
    transforms.ToTensor(),
])

# Create an instance of the SatelliteDataset
dataset = SatelliteDataset(
    image_folder=image_folder,
    label_folder=label_folder,
    image_transform=image_transform,
    label_transform=label_transform
)

# Create DataLoader for training and validation
batch_size = 3
train_dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
val_dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False)

# U-Net like segmentation model (adjusted to output 512x512)
class SegmentationCNN(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(SegmentationCNN, self).__init__()

        # Encoder (downsampling)
        self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=3, stride=1, padding=1)
        self.relu1 = nn.ReLU()
        self.maxpool1 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.conv2 = nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1)
        self.relu2 = nn.ReLU()
        self.maxpool2 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.conv3 = nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1)
        self.relu3 = nn.ReLU()
        self.maxpool3 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.conv4 = nn.Conv2d(256, 512, kernel_size=3, stride=1, padding=1)
        self.relu4 = nn.ReLU()

        # Decoder (upsampling)
        self.upconv1 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)
        self.conv5 = nn.Conv2d(512, 256, kernel_size=3, stride=1, padding=1)  # 256+256 from skip connection
        self.relu5 = nn.ReLU()

        self.upconv2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)
        self.conv6 = nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1)  # 128+128 from skip connection
        self.relu6 = nn.ReLU()

        self.upconv3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)
        self.conv7 = nn.Conv2d(128, out_channels, kernel_size=3, stride=1, padding=1)  # 64+64 from skip connection

    def forward(self, x):
        # Encoder pass
        x1 = self.relu1(self.conv1(x))
        x = self.maxpool1(x1)

        x2 = self.relu2(self.conv2(x))
        x = self.maxpool2(x2)

        x3 = self.relu3(self.conv3(x))
        x = self.maxpool3(x3)

        x = self.relu4(self.conv4(x))

        # Decoder pass with skip connections
        x = self.upconv1(x)
        x = torch.cat([x, x3], dim=1)
        x = self.relu5(self.conv5(x))

        x = self.upconv2(x)
        x = torch.cat([x, x2], dim=1)
        x = self.relu6(self.conv6(x))

        x = self.upconv3(x)
        x = torch.cat([x, x1], dim=1)
        x = self.conv7(x)

        return x


# Training loop
def train(model, dataloader, loss_fn, optimizer, device):
    model.train()
    running_loss = 0.0

    for images, labels in dataloader:
        images = images.to(device)
        labels = labels.to(device)

        optimizer.zero_grad()

        outputs = model(images)
        labels = labels.long()

        loss = loss_fn(outputs, labels)
        loss.backward()
        optimizer.step()

        running_loss += loss.item() * images.size(0)

    epoch_loss = running_loss / len(dataloader.dataset)
    return epoch_loss


# Validation
def evaluate(model, dataloader, loss_fn, device):
    model.eval()
    running_loss = 0.0

    with torch.no_grad():
        for images, labels in dataloader:
            images = images.to(device)
            labels = labels.to(device)

            outputs = model(images)
            labels = labels.long()

            loss = loss_fn(outputs, labels)
            running_loss += loss.item() * images.size(0)

    epoch_loss = running_loss / len(dataloader.dataset)
    return epoch_loss


# Perform training
def train_model(model, train_dataloader, val_dataloader, loss_fn, optimizer, device, num_epochs):
    best_val_loss = float('inf')
    best_model_weights = None
    for epoch in range(num_epochs):
        print(f"Epoch {epoch+1}/{num_epochs}")

        train_loss = train(model, train_dataloader, loss_fn, optimizer, device)
        print(f"Train Loss: {train_loss:.4f}")

        val_loss = evaluate(model, val_dataloader, loss_fn, device)
        print(f"Val Loss: {val_loss:.4f}")

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_model_weights = model.state_dict()

    model.load_state_dict(best_model_weights)
    return model


# Initialize components
model = SegmentationCNN(in_channels=3, out_channels=7)
device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")  # 优先使用Mac的MPS加速
model.to(device)

loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
num_epochs = 10

# Start training
model = train_model(model, train_dataloader, val_dataloader, loss_fn, optimizer, device, num_epochs)

验证说明

  1. 修复后标签保持单通道,batch_size=3时标签形状为[3,512,512],输入图像形状为[3,3,512,512],模型输出形状为[3,7,512,512]
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 23:18:47