如何结合过采样欠采样与SMOTE实现PyTorch多分类样本均衡
PyTorch下结合过采样与欠采样的多分类不均衡数据集处理方案
核心思路
你的目标是把所有类别样本量统一到40000条,刚好可以对应三类处理逻辑:
- 样本量超过40000的类别0(250000条)、类别1(48000条):执行随机欠采样,降到40000条
- 样本量等于40000的类别2:直接保留原始样本即可
- 样本量不足40000的类别3(38000条)、类别4(35000条)、类别5(7000条):用SMOTE算法做过采样,补到40000条
下面提供两种常用的落地实现方案:
方案1:离线预处理后加载(推荐全量数据可载入内存的场景)
一次性完成采样再喂入PyTorch DataLoader,实现简单、训练速度快。
需要提前安装imbalanced-learn库依赖。
import numpy as np from imblearn.over_sampling import SMOTE from imblearn.under_sampling import RandomUnderSampler from imblearn.pipeline import Pipeline import torch from torch.utils.data import TensorDataset, DataLoader # 假设你已经加载好了全量特征X(形状[N, 特征维度])和标签y(形状[N,]) # 定义每个类别的目标样本量 target_num = 40000 sampling_strategy = {0: target_num, 1: target_num, 2: target_num, 3: target_num, 4: target_num, 5: target_num} # 拆分欠采样、过采样的目标类别,先做欠采样减少后续SMOTE的计算量 under_strategy = {k:v for k,v in sampling_strategy.items() if v < np.sum(y == k)} over_strategy = {k:v for k,v in sampling_strategy.items() if v > np.sum(y == k)} # 构建采样流水线 resample_pipeline = Pipeline([ ('under_sampler', RandomUnderSampler(sampling_strategy=under_strategy, random_state=42)), ('smote', SMOTE(sampling_strategy=over_strategy, random_state=42)) ]) # 执行采样 X_resampled, y_resampled = resample_pipeline.fit_resample(X, y) # 转成PyTorch格式构建加载器 X_tensor = torch.tensor(X_resampled, dtype=torch.float32) y_tensor = torch.tensor(y_resampled, dtype=torch.long) dataset = TensorDataset(X_tensor, y_tensor) dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
方案2:自定义Sampler动态采样(推荐避免过拟合、大内存需求场景)
提前生成小类的SMOTE扩充样本池,每轮训练从每个类别样本池里随机抽40000条样本组合成训练集,每轮的采样组合都不同,泛化性更好。
from torch.utils.data import Sampler from collections import defaultdict # 按类别拆分原始样本索引 class_indices = defaultdict(list) for idx, label in enumerate(y): class_indices[label].append(idx) # 提前对小类做SMOTE过采样生成扩充样本池 over_sampler = SMOTE(sampling_strategy={3:40000,4:40000,5:40000}, random_state=42) X_aug, y_aug = over_sampler.fit_resample(X, y) # 统计扩充后所有类别的样本索引 aug_class_indices = defaultdict(list) for idx, label in enumerate(y_aug): aug_class_indices[label].append(idx) # 自定义均衡采样器 class BalancedClassSampler(Sampler): def __init__(self, class_indices, per_class_num=40000): self.class_indices = class_indices self.per_class_num = per_class_num self.total_num = len(class_indices) * per_class_num def __iter__(self): sampled = [] for label in self.class_indices: # 从对应类别样本池里无放回采样指定数量 sampled_label = np.random.choice(self.class_indices[label], size=self.per_class_num, replace=False) sampled.extend(sampled_label) # 全局打乱后返回 np.random.shuffle(sampled) return iter(sampled.tolist()) def __len__(self): return self.total_num # 构建加载器 X_aug_tensor = torch.tensor(X_aug, dtype=torch.float32) y_aug_tensor = torch.tensor(y_aug, dtype=torch.long) aug_dataset = TensorDataset(X_aug_tensor, y_aug_tensor) dataloader = DataLoader(aug_dataset, batch_size=64, sampler=BalancedClassSampler(aug_class_indices))
注意事项
- 如果你处理的是图像、文本这类非结构化数据,不要直接用普通SMOTE插值,插值生成的非结构化样本没有实际语义,建议改用生成式模型或者针对特定模态的过采样方法。
- 如果担心随机欠采样丢失大类的重要信息,可以把
RandomUnderSampler替换为NearMiss等带加权逻辑的欠采样算法,无需修改其他代码逻辑。
内容的提问来源于stack exchange,提问作者Shorouk Adel
相关产品推荐
相关产品推荐

