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

如何在PyTorch中对random_split后的训练集单独做数据增强?

实现仅训练集应用数据增强的方案

针对你的需求,核心思路是让自定义数据集支持动态传入transform,然后对划分后的训练集和验证集分别绑定不同的预处理规则,以下是两种可行的实现方式:

方法一:用包装类包裹划分后的子集

这种方式不需要修改原自定义Dataset的使用逻辑,通过一个轻量的包装类给不同子集绑定transform:

1. 自定义Dataset(支持transform参数)

确保你的Dataset类在初始化时接受transform和target_transform参数,并在__getitem__中应用:

import torch
from torch.utils.data import Dataset, random_split, DataLoader
from torchvision import transforms

class CustomDataset(Dataset):
    def __init__(self, data_paths, transform=None, target_transform=None):
        self.data_paths = data_paths  # 替换为你的数据路径/数据源
        self.transform = transform
        self.target_transform = target_transform

    def __len__(self):
        return len(self.data_paths)

    def __getitem__(self, idx):
        # 替换为你实际的数据和标签加载逻辑
        data = self._load_data(self.data_paths[idx])
        label = self._load_label(self.data_paths[idx])

        # 应用transform
        if self.transform:
            data = self.transform(data)
        if self.target_transform:
            label = self.target_transform(label)
        
        return data, label

    def _load_data(self, path):
        # 示例:读取图片(需根据你的数据类型调整)
        from PIL import Image
        return Image.open(path).convert('RGB')
    
    def _load_label(self, path):
        # 示例:从文件名提取标签(需根据你的标签逻辑调整)
        return int(path.split('_')[-1].split('.')[0])

2. 定义训练/验证集的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])
])

# 验证集仅做基础预处理
valid_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])
])

3. 划分数据集并绑定transform

创建一个包装类,给划分后的Subset绑定对应的transform:

class TransformSubset(Dataset):
    def __init__(self, subset, transform=None, target_transform=None):
        self.subset = subset
        self.transform = transform
        self.target_transform = target_transform

    def __len__(self):
        return len(self.subset)

    def __getitem__(self, idx):
        data, label = self.subset[idx]
        if self.transform:
            data = self.transform(data)
        if self.target_transform:
            label = self.target_transform(label)
        return data, label

# 第一步:创建完整数据集(暂不指定transform)
full_dataset = CustomDataset(all_data_paths)  # all_data_paths是你的所有数据路径列表

# 第二步:划分训练集和验证集
train_size = int(0.8 * len(full_dataset))
valid_size = len(full_dataset) - train_size
train_subset, valid_subset = random_split(full_dataset, [train_size, valid_size], generator=torch.Generator().manual_seed(42))

# 第三步:给子集绑定对应的transform
train_data = TransformSubset(train_subset, transform=train_transform)
valid_data = TransformSubset(valid_subset, transform=valid_transform)

4. 传入DataLoader使用

train_loader = DataLoader(train_data, batch_size=32, shuffle=True)
valid_loader = DataLoader(valid_data, batch_size=32, shuffle=False)

方法二:创建不同transform的数据集实例

这种方式是基于相同的索引划分,创建两个独立的Dataset实例(分别带训练和验证transform):

# 创建两个数据集实例,分别绑定训练和验证transform
full_train_dataset = CustomDataset(all_data_paths, transform=train_transform)
full_valid_dataset = CustomDataset(all_data_paths, transform=valid_transform)

# 生成统一的划分索引(保证训练/验证集数据一致)
generator = torch.Generator().manual_seed(42)
train_indices, valid_indices = random_split(range(len(full_train_dataset)), [train_size, valid_size], generator=generator)

# 根据索引创建子集
train_data = torch.utils.data.Subset(full_train_dataset, train_indices)
valid_data = torch.utils.data.Subset(full_valid_dataset, valid_indices)

# 后续传入DataLoader的方式和方法一一致

两种方法都能实现需求:方法一更节省内存(共享原数据集的数据源),方法二逻辑更直观。根据你的实际场景选择即可。

内容的提问来源于stack exchange,提问作者alice

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 11:05:23