PyTorch中如何为拆分后的训练集和测试集分别应用不同Transform?
为拆分后的训练/测试集设置不同Transform的方法
当你用random_split拆分ImageFolder数据集后,直接修改原数据集的transform会同时影响训练和测试集——因为两者都是原数据集的Subset,共享原有的transform配置。下面是最稳妥的解决方法:
步骤1:创建不带初始Transform的数据集
先不给ImageFolder设置transform,避免后续拆分后互相干扰:
import torchvision.transforms as transforms from torchvision.datasets import ImageFolder from torch.utils.data import random_split, Dataset init_dataset = ImageFolder(root=path_to_data) # 暂不指定transform train_data, test_data = random_split(init_dataset, [400, 116])
步骤2:编写Subset包装类
这个类的作用是给每个Subset绑定独立的transform,在获取数据时自动应用:
class TransformedSubset(Dataset): def __init__(self, subset, transform=None): self.subset = subset self.transform = transform def __getitem__(self, idx): img, label = self.subset[idx] if self.transform: img = self.transform(img) return img, label def __len__(self): return len(self.subset)
步骤3:定义训练/测试各自的Transform
根据需求分别设置策略(训练集通常加数据增强,测试集只做基础预处理):
# 训练集Transform(带数据增强) train_transform = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 测试集Transform(仅基础预处理) test_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])
步骤4:绑定Transform与拆分后的数据集
把拆分好的Subset和对应的Transform绑定,得到最终可用的数据集:
train_dataset = TransformedSubset(train_data, train_transform) test_dataset = TransformedSubset(test_data, test_transform)
处理完成后,你可以直接用这两个包装后的数据集创建DataLoader,进行后续的训练和测试流程。
内容的提问来源于stack exchange,提问作者saeid rasouli
相关产品推荐
相关产品推荐

