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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 16:57:20