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

如何在PyTorch中通过Subset/Dataloader实现训练与验证测试不同数据变换?

PyTorch实现训练/验证测试集差异化数据变换(基于Subset)

核心思路是基于同一个原始数据集,通过Subset划分训练/验证测试索引,再用自定义子集类为不同分区绑定对应的数据变换,避免创建多个独立数据集。

实现步骤

  1. 定义差异化的数据变换
    首先分别定义训练集的增强变换,以及验证/测试集的无增强变换:
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))
])
  1. 自定义带变换的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
  1. 创建原始数据集并划分索引
    原始数据集仅负责加载数据,不预设变换:
# 原始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))
  1. 创建带变换的子集与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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 08:35:25