PyTorch中基于批次的数据增强解决分类任务类别不平衡问题
针对类别不平衡的动态图像增强解决方案
核心问题分析
你之前对全数据集统一应用固定变换,既丢失了原始样本信息,又没针对性解决少数类数据不足的问题——多数类样本被无意义变换稀释,少数类的扩充也缺乏多样性,自然无法提升模型性能。
最优解决思路:少数类动态增强 + 加权采样结合
直接在数据加载阶段仅对少数类样本动态应用增强变换,同时保留原始样本特征,再配合你已有的WeightedRandomSampler保证少数类的采样占比,既能扩充少数类的数据多样性,又不会干扰多数类的学习。
具体代码修改步骤
1. 自定义条件增强变换类
这个类会根据样本所属类别,决定是否应用增强:
class ConditionalAugmentTransform: def __init__(self, base_transform, augment_transform, minority_classes): self.base_transform = base_transform # 基础预处理(裁剪、resize、转Tensor) self.augment_transform = augment_transform # 少数类专属增强 self.minority_classes = minority_classes # 预先确定的少数类列表 def __call__(self, img, target): # 先执行基础预处理 processed_img = self.base_transform(img) # 仅对少数类样本应用额外增强 if target in self.minority_classes: processed_img = self.augment_transform(processed_img) return processed_img, target
2. 修改get_data函数,整合动态增强
先统计类别分布确定少数类,再给训练集绑定条件变换,测试集保持基础预处理:
def get_data(size_image, root, batch_size, num_workers): # 定义基础预处理(无增强,用于所有样本) base_transform = transforms.Compose([ MaxCenterCrop(), transforms.Resize(size_image), transforms.ToTensor() ]) # 定义少数类增强变换(根据任务调整,比如植物分类适合这些操作) augment_transform = transforms.Compose([ transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(degrees=15), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2) ]) # 先加载无变换的训练集,统计类别样本数 trainset = Plantnet(root, 'images_train', transform=None) train_class_count = Counter(trainset.targets) # 确定少数类:样本数低于所有类平均值的一半(可根据实际情况调整阈值) avg_sample_count = sum(train_class_count.values()) / len(train_class_count) minority_classes = [cls for cls, cnt in train_class_count.items() if cnt < avg_sample_count * 0.5] # 给训练集绑定条件增强变换 trainset.transform = ConditionalAugmentTransform(base_transform, augment_transform, minority_classes) # 测试集只用基础预处理,保证评估一致性 testset = Plantnet(root, 'images_test', transform=base_transform) # 保留你原有的加权采样逻辑 ... # 你的原有代码:计算weights、创建WeightedRandomSampler等 trainloader = torch.utils.data.DataLoader(trainset, batch_size=batch_size, sampler=sampler, shuffle=False, num_workers=num_workers) testloader = torch.utils.data.DataLoader(testset, batch_size=batch_size, shuffle=False, num_workers=num_workers) return trainloader, testloader, dataset_attributes
3. 训练函数无需额外修改
你的train函数逻辑可以完全保留,数据加载器会自动返回经过条件增强的样本,训练过程中少数类样本每次都会有不同的增强结果,相当于动态扩充了数据集。
关键注意事项
- 增强变换要贴合任务:植物分类避免使用会破坏植物形态的极端变换(比如大角度旋转、随机裁剪到主体外)。
- 少数类阈值灵活调整:如果类别不平衡特别严重(比如某类样本数是多数类的1/10),可以把阈值降到平均值的1/3甚至更低。
- 验证增强效果:可以在验证集上测试不同增强组合的效果,选择最优的增强策略。
内容的提问来源于stack exchange,提问作者nicenoize
相关产品推荐
相关产品推荐

