能否使用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
相关产品推荐
相关产品推荐

