如何在PyTorch ImageFolder中不新增本地图片实现样本观测数扩增
PyTorch无需修改本地文件的数据集增广实现方案
你原有代码中图像变换是每次取样本时实时执行的,但ImageFolder的默认长度是本地文件夹内的原始图片总数,因此不会统计增广后的样本量。要将每个变换生成的版本作为独立样本,可通过以下两种方案实现:
方法1:自定义数据集包装类(适配随机变换、灵活调整增广倍数)
通过自定义数据集类包装原始ImageFolder,直接指定单张原始图片对应的增广样本数量,动态调整总样本量:
import torch import os from torch.utils.data import Dataset, DataLoader from torchvision import transforms, datasets class AugmentedDataset(Dataset): def __init__(self, original_dataset, aug_per_sample): self.original_dataset = original_dataset self.aug_per_sample = aug_per_sample # 单张原始图片生成的增广样本数 def __len__(self): # 总样本量 = 原始样本量 * 单样本增广倍数 return len(self.original_dataset) * self.aug_per_sample def __getitem__(self, idx): # 计算当前索引对应的原始样本索引 original_idx = idx // self.aug_per_sample # 每次调用都会基于随机变换生成不同的增广版本 img, label = self.original_dataset[original_idx] return img, label
调用示例(假设单张原始图生成5个增广版本):
trf = transforms.Compose([ transforms.ToTensor(), transforms.RandomRotation(degrees=45), transforms.Grayscale(num_output_channels=1), transforms.Normalize(0, 1), transforms.functional.invert ]) original_train_data = datasets.ImageFolder(root='./splitted_data/train', transform= trf) # 增广后总样本量为原始的5倍 augmented_train_data = AugmentedDataset(original_train_data, aug_per_sample=5) print(len(augmented_train_data)) # 输出为原始长度 * 5 train_loader = DataLoader(augmented_train_data, batch_size=32, shuffle=True, num_workers=os.cpu_count())
方法2:拼接多变换数据集(适配固定多变换组合对比)
如果需要测试不同固定变换组合的效果,可通过ConcatDataset拼接多个使用不同变换的ImageFolder实例:
from torch.utils.data import ConcatDataset # 定义不同的变换组合 trf1 = transforms.Compose([ transforms.ToTensor(), transforms.RandomRotation(degrees=45), transforms.Grayscale(num_output_channels=1), transforms.Normalize(0, 1), transforms.functional.invert ]) trf2 = transforms.Compose([ transforms.ToTensor(), transforms.RandomHorizontalFlip(p=1), transforms.Grayscale(num_output_channels=1), transforms.Normalize(0, 1), transforms.functional.invert ]) trf3 = transforms.Compose([ transforms.ToTensor(), transforms.RandomResizedCrop(size=224), transforms.Grayscale(num_output_channels=1), transforms.Normalize(0, 1), transforms.functional.invert ]) # 分别创建对应不同变换的数据集实例 ds1 = datasets.ImageFolder(root='./splitted_data/train', transform=trf1) ds2 = datasets.ImageFolder(root='./splitted_data/train', transform=trf2) ds3 = datasets.ImageFolder(root='./splitted_data/train', transform=trf3) # 拼接后总样本量为原始的3倍 combined_train_data = ConcatDataset([ds1, ds2, ds3]) train_loader = DataLoader(combined_train_data, batch_size=32, shuffle=True, num_workers=os.cpu_count())
注意事项
- 若使用带随机性的变换(如
RandomRotation),方法1足够满足需求,仅需调整aug_per_sample参数即可快速修改总样本量,适配调试需求 - 若需要对比不同变换组合的训练效果,选择方法2结构更清晰,新增变换仅需追加对应的数据集实例即可
- 两种方案所有增广操作均为实时生成,不会修改本地图片文件
内容的提问来源于stack exchange,提问作者user9102437
相关产品推荐
相关产品推荐

