PyTorch卫星图像分割中输入与目标batch_size不匹配问题求助
卫星图像分割代码中Batch Size不匹配问题排查与解决
问题描述
运行卫星图像7类分割代码时,出现如下错误:
ValueError: Expected input batch_size (3) to match target batch_size (9).
无论调整batch_size,目标batch_size始终是输入的3倍,无法定位问题根源。
问题根源
- 标签被错误应用图像归一化变换:代码中对标签和图像使用了相同的
transform(包含Normalize),导致单通道的类别掩码被转换成3通道张量。后续labels.view(-1, 512, 512)会将[3,3,512,512]的张量重塑为[9,512,512],直接导致目标batch_size变为输入的3倍。 - 模型输出与标签尺寸不匹配:模型经过三次
MaxPool2d(每次下采样2倍),输入512x512图像最终输出尺寸为64x64,但标签始终保持512x512,这会进一步引发维度不匹配问题。 - 冗余的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)
验证说明
- 修复后标签保持单通道,
batch_size=3时标签形状为[3,512,512],输入图像形状为[3,3,512,512],模型输出形状为[3,7,512,512]
相关产品推荐
相关产品推荐

