如何在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
相关产品推荐
相关产品推荐

