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

能否使用torchvision.transforms实现PyTorch语义分割任务的数据增强?

语义分割中用PyTorch transforms.compose()同时处理图像和标签的实现

当然可以实现,但PyTorch原生的大部分transforms组件仅针对单张图像设计,无法直接同时处理语义分割的标签掩码。要实现同步增强,有两种主流方案:纯PyTorch原生自定义transform,或是结合Albumentations库(更高效便捷)并包装为PyTorch兼容形式。以下是两种方案的完整示例:

方案一:纯PyTorch原生实现

通过自定义transform类,让每个增强操作同时作用于图像和标签,再用transforms.compose()组合流水线。

1. 导入依赖

import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
import numpy as np
from PIL import Image

2. 自定义语义分割数据集

class SegmentationDataset(Dataset):
    def __init__(self, image_paths, mask_paths, transform=None):
        self.image_paths = image_paths
        self.mask_paths = mask_paths
        self.transform = transform

    def __len__(self):
        return len(self.image_paths)

    def __getitem__(self, idx):
        # 读取图像和标签,转为numpy数组
        image = np.array(Image.open(self.image_paths[idx]).convert("RGB"))
        # 标签用int64类型,避免后续训练时的梯度计算问题
        mask = np.array(Image.open(self.mask_paths[idx]).convert("L"), dtype=np.int64)

        if self.transform:
            # 将图像和标签打包为字典传入transform
            augmented = self.transform({"image": image, "mask": mask})
            image = augmented["image"]
            mask = augmented["mask"]

        # 图像转为PyTorch张量(HWC -> CHW),标签直接转张量
        image = transforms.ToTensor()(image)
        mask = torch.from_numpy(mask)

        return image, mask

3. 定义兼容图像和标签的自定义transform

class CustomHorizontalFlip:
    def __init__(self, p=0.5):
        self.p = p

    def __call__(self, sample):
        image, mask = sample["image"], sample["mask"]
        if torch.rand(1) < self.p:
            # 对numpy数组进行水平翻转,同步作用于图像和标签
            image = np.fliplr(image)
            mask = np.fliplr(mask)
        return {"image": image, "mask": mask}

class CustomResize:
    def __init__(self, size):
        self.size = size

    def __call__(self, sample):
        image, mask = sample["image"], sample["mask"]
        # 图像用双线性插值,标签用最近邻插值防止类别模糊
        image = np.array(Image.fromarray(image).resize(self.size))
        mask = np.array(Image.fromarray(mask).resize(self.size, Image.NEAREST))
        return {"image": image, "mask": mask}

4. 组合transform并构建DataLoader

# 构建增强流水线
transform = transforms.Compose([
    CustomResize((256, 256)),
    CustomHorizontalFlip(p=0.5),
    # 可添加更多自定义transform,比如随机旋转、裁剪等
])

# 替换为你的实际图像和标签路径列表
image_paths = ["img_01.jpg", "img_02.jpg", "img_03.jpg"]
mask_paths = ["mask_01.png", "mask_02.png", "mask_03.png"]

# 创建数据集和DataLoader
dataset = SegmentationDataset(image_paths, mask_paths, transform=transform)
dataloader = DataLoader(dataset, batch_size=4, shuffle=True)

# 测试数据加载
for imgs, masks in dataloader:
    print(f"批量图像形状: {imgs.shape}, 批量标签形状: {masks.shape}")
    break

方案二:结合Albumentations与PyTorch transforms

Albumentations原生支持同时处理图像和标签,且增强操作更丰富、速度更快,只需将其包装为PyTorch兼容的形式即可结合transforms.compose()使用(或直接用Albumentations自己的Compose)。

1. 导入依赖

import torch
from torch.utils.data import Dataset, DataLoader
import numpy as np
from PIL import Image
import albumentations as A
from albumentations.pytorch import ToTensorV2

2. 自定义数据集(适配Albumentations)

class SegmentationDatasetAlb(Dataset):
    def __init__(self, image_paths, mask_paths, transform=None):
        self.image_paths = image_paths
        self.mask_paths = mask_paths
        self.transform = transform

    def __len__(self):
        return len(self.image_paths)

    def __getitem__(self, idx):
        image = np.array(Image.open(self.image_paths[idx]).convert("RGB"))
        mask = np.array(Image.open(self.mask_paths[idx]).convert("L"), dtype=np.int64)

        if self.transform:
            # Albumentations直接接收image和mask参数,返回增强后的结果
            augmented = self.transform(image=image, mask=mask)
            image = augmented["image"]
            mask = augmented["mask"]

        return image, mask

3. 定义增强流水线并构建DataLoader

# 用Albumentations定义增强,自动同步处理图像和标签
alb_transform = A.Compose([
    A.Resize(height=256, width=256),
    A.HorizontalFlip(p=0.5),
    A.RandomRotate90(p=0.3),
    A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),
    ToTensorV2()  # 自动将图像转为CHW张量,标签转为张量
])

# 创建数据集和DataLoader
dataset_alb = SegmentationDatasetAlb(image_paths, mask_paths, transform=alb_transform)
dataloader_alb = DataLoader(dataset_alb, batch_size=4, shuffle=True)

# 测试
for imgs, masks in dataloader_alb:
    print(f"批量图像形状: {imgs.shape}, 批量标签形状: {masks.shape}")
    break

说明:方案二更推荐,因为Albumentations提供了大量语义分割专用的增强操作(如RandomCrop带标签对齐、CoarseDropout对标签无影响等),且底层基于numpy优化,运行速度比纯PyTorch实现更快。

内容的提问来源于stack exchange,提问作者KFK

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 18:57:49