PyTorch曝光检测实现报错:目标与输入尺寸不匹配的解决方法
图像曝光校正训练尺寸不匹配ValueError解决方案
问题现象
训练基于预训练DenseNet121的图像曝光校正模型时,触发ValueError:
目标尺寸(torch.Size([16, 1, 224, 224]))必须与输入尺寸(torch.Size([16, 1]))一致
根因分析
- 数据集输出错误:
ImagePairDataset的__getitem__方法中,返回target_img[0]会将(C, H, W)格式的目标图像切片为(H, W)张量,批量后经unsqueeze(1)处理变成(16,1,224,224),与模型输出的(16,1)维度完全不匹配。 - 任务类型混淆:当前模型被改造成二分类任务(输出维度1),但图像曝光校正属于像素级图像回归任务,需要输出与输入尺寸一致的图像,任务类型不匹配导致尺寸冲突。
解决方案
1. 修正数据集的目标输出
移除__getitem__中target_img[0]的切片操作,返回完整目标图像张量:
def __getitem__(self, index): train_img_name = self.input_filenames[index].split("-")[0] train_img = Image.open(os.path.join(self.input_dir, self.input_filenames[index])) for i in range(len(self.target_name_list)): if train_img_name == self.target_name_list[i]: target_img = Image.open(os.path.join(self.target_dir, self.target_filenames[i])) break if self.transform: train_img = self.transform(train_img) target_img = self.transform(target_img) # 移除[0],返回完整的目标图像张量 return train_img, target_img
2. 重构模型为图像到图像回归网络
将原分类用DenseNet121修改为输出与输入尺寸一致的回归模型:
import torch import torch.nn as nn from torchvision.models import densenet121 image_size = 224 model = densenet121(pretrained=True) # 替换分类器为转置卷积层,恢复特征图到输入图像尺寸 model.classifier = nn.Sequential( nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2), nn.ReLU(inplace=True), nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2), nn.ReLU(inplace=True), nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2), nn.ReLU(inplace=True), nn.ConvTranspose2d(128, 3, kernel_size=2, stride=2), # 3通道RGB输入对应3通道输出 nn.Sigmoid() # 将输出像素值归一化到0-1范围 ) # 若为单通道灰度图像,将最后一层改为nn.ConvTranspose2d(128, 1, kernel_size=2, stride=2) model.to(device)
3. 更换损失函数
图像回归任务使用MSE损失更合适,替换原BCEWithLogitsLoss:
criterion = nn.MSELoss()
4. 修正训练循环的目标处理
移除targets.unsqueeze(1),直接用原始目标张量计算损失:
for inputs, targets in train_dataloader: inputs = inputs.to(device) targets = targets.to(device) optimiser.zero_grad() outputs = model(inputs) # 直接计算损失,无需修改目标尺寸 loss = criterion(outputs, targets.float()) loss.backward() optimiser.step() running_loss += loss.item() * inputs.size(0)
额外优化:数据集配对逻辑
将原循环配对改为字典映射,提升效率:
def __init__(self, input_dir, target_dir, transform=None): self.input_dir = input_dir self.target_dir = target_dir self.transform = transform # 构建前缀到目标文件名的映射字典 self.target_map = {} for fname in os.listdir(target_dir): prefix = fname.split("-")[0] self.target_map[prefix] = fname self.input_filenames = os.listdir(input_dir) def __getitem__(self, index): input_fname = self.input_filenames[index] prefix = input_fname.split("-")[0] train_img = Image.open(os.path.join(self.input_dir, input_fname)) # 通过映射直接获取目标文件,避免循环遍历 target_fname = self.target_map[prefix] target_img = Image.open(os.path.join(self.target_dir, target_fname)) if self.transform: train_img = self.transform(train_img) target_img = self.transform(target_img) return train_img, target_img
内容的提问来源于stack exchange,提问作者Yadhu B
相关产品推荐
相关产品推荐

