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

数据增强时图像与掩码无法同步应用变换的问题求助

问题根源

你当前代码的核心问题是分别对图像和掩码调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 05:50:58