如何在PyTorch中通过Subset/Dataloader实现训练与验证测试不同数据变换?
PyTorch实现训练/验证测试集差异化数据变换(基于Subset)
核心思路是基于同一个原始数据集,通过Subset划分训练/验证测试索引,再用自定义子集类为不同分区绑定对应的数据变换,避免创建多个独立数据集。
实现步骤
- 定义差异化的数据变换
首先分别定义训练集的增强变换,以及验证/测试集的无增强变换:
import torch from torch.utils.data import Dataset, Subset, DataLoader from torchvision import transforms from torchvision.datasets import CIFAR10 # 训练集增强变换 train_transform = transforms.Compose([ transforms.RandomCrop(32, padding=4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)) ]) # 验证/测试集无增强变换 val_test_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)) ])
- 自定义带变换的Subset子类
继承Subset类,添加变换属性,在获取样本时自动应用对应变换:
class TransformSubset(Subset): def __init__(self, dataset, indices, transform=None): super().__init__(dataset, indices) self.transform = transform def __getitem__(self, idx): # 获取原始样本(图片+标签) img, label = super().__getitem__(idx) # 应用当前子集对应的变换 if self.transform is not None: img = self.transform(img) return img, label
- 创建原始数据集并划分索引
原始数据集仅负责加载数据,不预设变换:
# 原始CIFAR10数据集,仅加载数据,不做变换 original_dataset = CIFAR10(root='./data', train=True, download=True, transform=None) # 手动划分训练/验证索引(示例:前45000为训练,后5000为验证) train_idx = list(range(45000)) val_idx = list(range(45000, 50000))
- 创建带变换的子集与DataLoader
为训练子集绑定增强变换,验证/测试子集绑定无增强变换,再传入DataLoader:
# 训练子集(带数据增强) train_subset = TransformSubset(original_dataset, train_idx, transform=train_transform) # 验证子集(无数据增强) val_subset = TransformSubset(original_dataset, val_idx, transform=val_test_transform) # 创建数据加载器 train_loader = DataLoader(train_subset, batch_size=64, shuffle=True, num_workers=2) val_loader = DataLoader(val_subset, batch_size=64, shuffle=False, num_workers=2)
测试集处理
对于测试集,只需基于原始测试数据集创建TransformSubset并绑定val_test_transform即可:
test_dataset = CIFAR10(root='./data', train=False, download=True, transform=None) test_subset = TransformSubset(test_dataset, list(range(len(test_dataset))), transform=val_test_transform) test_loader = DataLoader(test_subset, batch_size=64, shuffle=False, num_workers=2)
这种方式全程仅维护一个原始数据集,通过Subset实现数据分区,同时为不同分区灵活配置变换,完全符合你的需求。
内容的提问来源于stack exchange,提问作者ChaosPredictor
相关产品推荐
相关产品推荐

