如何在PyTorch中仅对过采样生成的新数据应用数据增强?
解决方案
1. 仅对过采样的重复数据应用图像增强
原代码将增强变换直接作用于全部过采样数据集,导致原始样本也被过度增强,结合过采样的重复样本极易引发过拟合。要实现仅对重复采样的样本做增强,需自定义Dataset区分原始样本与重复样本:
实现步骤
- 先统计每个类别的样本数量,计算每个原始样本需要被采样的总次数
- 自定义
OversampleAugDataset,在获取样本时判断是否为重复采样:是则应用增强变换,否则仅做基础格式转换
代码示例
首先定义两种转换规则:
# 基础转换:仅转换为Tensor,用于原始样本 transform_base = transforms.Compose([ transforms.ToTensor(), ]) # 增强转换:仅用于过采样生成的重复样本 transform_aug = transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.RandomResizedCrop(size=150, scale=(0.8, 1.0)), transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4), transforms.RandomRotation(degrees=15), transforms.ToTensor(), ])
自定义数据集类:
class OversampleAugDataset(data.Dataset): def __init__(self, original_dataset, sample_counts): self.original_dataset = original_dataset self.samples = [] # 构建采样列表:记录每个样本的原始索引及是否需要增强 for idx in range(len(original_dataset)): count = sample_counts[idx] # 第一次采样为原始样本,不增强 self.samples.append((idx, False)) # 后续采样为重复样本,需增强 self.samples.extend([(idx, True) for _ in range(count - 1)]) def __len__(self): return len(self.samples) def __getitem__(self, idx): orig_idx, need_aug = self.samples[idx] img, label = self.original_dataset[orig_idx] return (transform_aug(img), label) if need_aug else (transform_base(img), label)
生成采样次数并创建DataLoader:
import numpy as np # 批量获取所有标签(优化后的方式见下文) labels = np.array([label for _, label in train_dataset]) class_counts = np.bincount(labels) max_class_size = class_counts.max() # 计算每个原始样本的总采样次数 sample_counts = [int(max_class_size / class_counts[label]) for label in labels] # 创建数据集与加载器 aug_oversample_dataset = OversampleAugDataset(train_dataset, sample_counts) train_loader = data.DataLoader(aug_oversample_dataset, batch_size=64, shuffle=True)
2. 优化类别索引提取效率
原循环遍历20000张图片逐个获取标签的方式效率极低,可通过numpy向量化操作批量处理:
方法1:批量提取标签并分组
import numpy as np # 一次性提取所有标签,替代循环遍历 labels = np.array([label for _, label in train_dataset]) # 按类别快速分组索引 class_indices = [np.where(labels == c)[0].tolist() for c in range(7)]
方法2:直接访问数据集的标签属性(若自定义Dataset时已存储)
如果你的train_dataset自定义时将标签存在了self.labels属性中,可直接调用:
import numpy as np labels = np.array(train_dataset.labels) class_indices = [np.where(labels == c)[0].tolist() for c in range(7)]
numpy的向量化操作效率远高于Python循环,能大幅缩短索引提取时间。
预期效果
修改后,原始样本仅做基础格式转换,仅过采样的重复样本被增强,既缓解了类别不平衡问题,又避免了原始样本过度增强导致的过拟合,训练准确率不会过早达到1.0,验证准确率也会逐步提升。
内容的提问来源于stack exchange,提问作者joebob
相关产品推荐
相关产品推荐

