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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 08:12:44