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

如何结合过采样欠采样与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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 03:48:01