PyTorch训练Camus数据集:图像尺寸不一致致批量训练异常的解决求助
解决Camus数据集图像尺寸统一问题
一、必须分离图像与掩码
图像和掩码是一一对应的标注对,预处理时必须同步执行尺寸变换操作,否则掩码会与图像错位,完全失去标注价值。所以第一步要确保加载时将图像和掩码分开处理,再用相同的尺寸逻辑同步转换。
二、两种常用的尺寸统一方案
根据你的需求选择合适的方案:
方案1:直接缩放(Resize)
将所有图像和掩码强制缩放到固定尺寸(比如(256,256)),优点是保留完整图像内容,实现简单;缺点是会带来一定比例拉伸变形,适合对图像比例要求不高的任务。
方案2:缩放+中心裁剪
先将图像缩放到短边匹配目标尺寸,再对长边进行中心裁剪,最大程度保留图像核心区域(Camus的心脏区域通常在中心),减少变形影响。
三、PyTorch具体实现代码
以下是基于自定义Dataset类的完整实现,以直接缩放为例:
import torch from torch.utils.data import Dataset from torchvision import transforms import cv2 import numpy as np class CamusDataset(Dataset): def __init__(self, sample_paths, target_size=(256, 256)): # sample_paths是列表,每个元素存储(图像路径, 掩码路径) self.sample_paths = sample_paths self.target_size = target_size # 图像变换:单通道灰度图,用双线性插值缩放 self.img_transform = transforms.Compose([ transforms.ToTensor(), transforms.Resize(target_size, antialias=True) ]) # 掩码变换:掩码是标签,必须用最近邻插值避免出现非整数标签值 self.mask_transform = transforms.Compose([ transforms.ToTensor(), transforms.Resize(target_size, interpolation=transforms.InterpolationMode.NEAREST) ]) def __len__(self): return len(self.sample_paths) def __getitem__(self, idx): img_path, mask_path = self.sample_paths[idx] # 加载单通道灰度图像 img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img = np.expand_dims(img, axis=-1) # 转为(H,W,1)格式,适配ToTensor # 加载掩码(同样是单通道) mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) mask = np.expand_dims(mask, axis=-1) # 同步执行尺寸变换 img_tensor = self.img_transform(img) mask_tensor = self.mask_transform(mask) return img_tensor, mask_tensor
如果要改用缩放+中心裁剪方案,只需替换变换逻辑:
# 自定义缩放+中心裁剪变换 class ResizeCenterCrop: def __init__(self, target_size): self.target_size = target_size def __call__(self, tensor): h, w = tensor.shape[-2:] target_h, target_w = self.target_size # 计算缩放比例,确保短边匹配目标尺寸 scale = max(target_h / h, target_w / w) new_h, new_w = int(h * scale), int(w * scale) # 先缩放 resized = transforms.Resize((new_h, new_w))(tensor) # 再中心裁剪 return transforms.CenterCrop(self.target_size)(resized) # 更新变换组合 self.img_transform = transforms.Compose([ transforms.ToTensor(), ResizeCenterCrop(self.target_size) ]) self.mask_transform = transforms.Compose([ transforms.ToTensor(), ResizeCenterCrop(self.target_size), transforms.Resize(self.target_size, interpolation=transforms.InterpolationMode.NEAREST) ])
四、批量训练验证
实例化数据集后,用DataLoader加载即可得到固定尺寸的批量张量:
# 假设已整理好样本路径列表sample_list dataset = CamusDataset(sample_list, target_size=(256,256)) dataloader = torch.utils.data.DataLoader(dataset, batch_size=8, shuffle=True) # 验证输出尺寸 for imgs, masks in dataloader: print(imgs.shape) # 输出torch.Size([8, 1, 256, 256]) print(masks.shape) # 输出torch.Size([8, 1, 256, 256]) break
内容的提问来源于stack exchange,提问作者Dimgai
相关产品推荐
相关产品推荐

