PyTorch中如何仅对特定类别的图像应用transform数据增强
PyTorch按类别触发特定数据增强的实现方案
PyTorch默认传入Dataset的transform参数会对所有样本无差别生效,要实现「仅特定类别应用指定变换、其他类别不应用/应用另一套变换」的需求,核心是把变换判断逻辑下沉到单样本读取环节,根据样本标签匹配对应的增强规则,常见实现方式有两种:
方法1:自定义Dataset子类(最灵活,适合自定义数据集)
直接重写数据集的__getitem__方法,在读取单张图片、拿到标签后,按标签选择提前定义好的对应变换流程即可。以猫狗分类、仅对猫类做增强的场景为例,代码实现如下:
import os from PIL import Image from torch.utils.data import Dataset from torchvision import transforms class CatDogDataset(Dataset): def __init__(self, data_root): # 类别映射,猫对应标签0,狗对应标签1 self.cls_map = {"cat": 0, "dog": 1} self.samples = [] # 遍历目录收集所有图片路径和对应标签 for cls_name, label in self.cls_map.items(): cls_folder = os.path.join(data_root, cls_name) for file_name in os.listdir(cls_folder): if file_name.lower().endswith((".jpg", ".jpeg", ".png")): self.samples.append((os.path.join(cls_folder, file_name), label)) # 定义猫类专属增强:包含随机翻转、颜色抖动等增强操作 self.cat_aug = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(degrees=15), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 定义狗类预处理:仅做尺寸调整、张量化、归一化,无额外增强 self.dog_preprocess = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] img = Image.open(img_path).convert("RGB") # 核心判断逻辑:按标签选择对应变换 if label == self.cls_map["cat"]: img = self.cat_aug(img) else: img = self.dog_preprocess(img) return img, label
方法2:变换包装器(适合内置数据集,无需重写整个Dataset)
如果使用CIFAR-10、FashionMNIST等PyTorch内置数据集,不想重写完整Dataset逻辑,可以写一个包装类包裹原始数据集,在包装层实现按类别选变换的逻辑:
class ClassSpecificTransform: def __init__(self, base_dataset, cls_transform_dict, default_transform=None): """ Args: base_dataset: 原始数据集,需返回(图像, 标签)格式,注意要把原始数据集的transform参数设为None cls_transform_dict: 类别-变换映射字典,key为类别标签值,value为对应要应用的transform default_transform: 未匹配到对应类别时使用的默认变换 """ self.base_dataset = base_dataset self.cls_transform_dict = cls_transform_dict self.default_transform = default_transform def __len__(self): return len(self.base_dataset) def __getitem__(self, idx): img, label = self.base_dataset[idx] if label in self.cls_transform_dict: img = self.cls_transform_dict[label](img) elif self.default_transform is not None: img = self.default_transform(img) return img, label
调用示例:
from torchvision.datasets import CIFAR10 # 初始化原始数据集,注意transform设为None base_set = CIFAR10(root="./data", train=True, download=True, transform=None) # 假设仅对标签为3的类别(猫)做增强,其余类别用基础预处理 cls_transform = { 3: transforms.Compose([transforms.RandomHorizontalFlip(), transforms.ToTensor()]) } default_trans = transforms.Compose([transforms.ToTensor()]) # 包装得到最终数据集 train_set = ClassSpecificTransform(base_set, cls_transform, default_trans)
注意事项
- 不要给原始数据集设置全局
transform参数,否则所有样本会先经过全局无差别变换,再进入类别判断逻辑,无法达到差异化增强的效果 - 不要在初始化阶段提前对所有图片做变换并缓存,否则随机增强的结果会被固定,每个epoch拿到的增强样本完全一致,失去数据增强的随机性意义
- 该逻辑可以灵活扩展,比如给不同类别配置不同强度的增强、过滤掉会生成无效样本的变换组合,适配各类业务场景的约束。
内容的提问来源于stack exchange,提问作者Paradom
相关产品推荐
相关产品推荐

