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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 20:10:27