数据增强时图像与掩码无法同步应用变换的问题求助
问题根源
你当前代码的核心问题是分别对图像和掩码调用transform,torchvision的随机变换(如RandomRotation、RandomFlip)每次调用都会重新生成随机参数(比如旋转角度、是否翻转),导致图像和掩码的变换不同步。
解决方案1:使用TorchVision v2统一变换(推荐)
TorchVision v2原生支持同时对图像和语义掩码执行同步变换,只需将两者打包传入transform即可,变换会自动复用同一组随机参数。
修改Dataset的__getitem__方法
def __getitem__(self, index): dict_path = os.path.join(self.dict_dir, self.data[index]) patient_dict = torch.load(dict_path) image = patient_dict['image'].unsqueeze(0) mass_mask = patient_dict['mass_mask'].unsqueeze(0) mass_mask[mass_mask > 1.0] = 1.0 if self.transform is not None: # 将图像和掩码作为元组传入,torchvision v2会同步应用变换 image, mass_mask = self.transform(image, mass_mask) return image, mass_mask
调整Transform参数(适配掩码填充)
掩码是分割任务的标签,背景填充值应该为0(而非图像的255),TorchVision v2支持为不同类型输入指定不同填充值:
train_transform = T.Compose( [ # fill参数传入元组:(图像填充值, 掩码填充值) T.RandomRotation(degrees=35, expand=True, fill=(255.0, 0.0)), T.RandomHorizontalFlip(p=0.5), T.RandomVerticalFlip(p=0.5), ] )
解决方案2:自定义同步变换类(兼容旧版TorchVision)
如果使用TorchVision v1,可以自定义变换类,提前采样一次随机参数,再同步应用到图像和掩码:
自定义变换类
import random from torchvision.transforms import functional as F class SyncedSegTransform: def __call__(self, image, mask): # 提前采样所有随机参数 rot_degree = random.uniform(-35, 35) do_hflip = random.random() < 0.5 do_vflip = random.random() < 0.5 # 同步应用旋转 image = F.rotate(image, rot_degree, expand=True, fill=255.0) mask = F.rotate(mask, rot_degree, expand=True, fill=0.0) # 同步应用水平翻转 if do_hflip: image = F.hflip(image) mask = F.hflip(mask) # 同步应用垂直翻转 if do_vflip: image = F.vflip(image) mask = F.vflip(mask) return image, mask
在Dataset中使用
def __getitem__(self, index): # ... 加载数据代码不变 ... if self.transform is not None: image, mass_mask = self.transform(image, mass_mask) return image, mass_mask
定义Transform
train_transform = SyncedSegTransform()
解决方案3:使用Albumentations(专业分割任务增强库)
Albumentations专为图像分割设计,天然支持图像与掩码的同步变换,只需按以下方式修改:
修改Dataset
import albumentations as A from albumentations.pytorch import ToTensorV2 import cv2 class INBreastDataset2012(Dataset): def __init__(self, dict_dir, transform=None): self.dict_dir = dict_dir self.data = os.listdir(self.dict_dir) self.transform = transform def __len__(self): return len(self.data) def __getitem__(self, index): dict_path = os.path.join(self.dict_dir, self.data[index]) patient_dict = torch.load(dict_path) # 转换为Albumentations要求的HWC格式numpy数组 image = patient_dict['image'].numpy()[..., None] # (H,W) -> (H,W,1) mass_mask = patient_dict['mass_mask'].numpy() mass_mask[mass_mask > 1.0] = 1.0 mass_mask = mass_mask[..., None] # (H,W) -> (H,W,1) if self.transform is not None: # 同步变换图像和掩码 transformed = self.transform(image=image, mask=mass_mask) image = transformed['image'] mass_mask = transformed['mask'] # 转换为PyTorch要求的CHW格式tensor image = image.permute(2, 0, 1) mass_mask = mass_mask.permute(2, 0, 1) return image, mass_mask
定义Albumentations Transform
train_transform = A.Compose([ A.RandomRotate(limit=35, p=1.0, border_mode=cv2.BORDER_CONSTANT, value=255, mask_value=0), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), ToTensorV2(), ])
内容的提问来源于stack exchange,提问作者GASTON DANIEL BAZAN
相关产品推荐
相关产品推荐

