使用Pandas和Sklearn进行数据集上采样:保持多类别比例的需求
问题需求
我的数据集存在严重类别不平衡:relevance列两类记录数分别为190条和14810条。尝试对少数类上采样后,发现另一列class的类别比例失衡(原各类别均为1000条左右)。现在需要实现:将relevance列两类数量拉平的同时,保持class列的类别比例为1:1:1。
现有上采样代码
# 创建少数类数据集 df_minority = df[df['relevance'] == 1] # 创建其余类别数据集 df_rest = df[df['relevance'] != 1] # 上采样少数类 df_1_upsampled = resample(df_minority, random_state=SEED, n_samples=14810, replace=True) # 合并上采样后的数据集 df_upsampled = pd.concat([df_1_upsampled, df_rest])
示例数据集
relevance class 2 3 4 5 1 A 40 24 11 50 1 A 60 20 19 60 0 C 15 57 15 60 0 B 12 50 15 43 0 B 90 8 32 80 0 C 74 8 21 34
解决方案
要同时满足两个约束,需要以relevance和class的组合为分层依据进行分组采样,确保上/下采样后两个列的分布都符合要求。
核心思路
- 设定目标:
relevance的两类记录数均为14810条,同时每个relevance分组下的3个class类别数量相等(即每个(relevance, class)组合的目标数量为14810 // 3 ≈ 4937,总数可微调保证一致)。 - 对少数类(
relevance=1)按class分组单独上采样,每个子组达到目标数量。 - 对多数类(
relevance=0)同样按class分组调整:原数量超过目标则下采样,不足则上采样,确保每个子组数量达标。
实现代码
from sklearn.utils import resample import pandas as pd SEED = 42 # 设定目标数量 target_relevance_total = 14810 target_per_class = target_relevance_total // 3 # 处理少数类:relevance=1 minority_groups = [] for cls in df[df['relevance'] == 1]['class'].unique(): subset = df[(df['relevance'] == 1) & (df['class'] == cls)] # 上采样至目标数量 upsampled = resample(subset, random_state=SEED, n_samples=target_per_class, replace=True) minority_groups.append(upsampled) minority_processed = pd.concat(minority_groups) # 处理多数类:relevance=0 majority_groups = [] for cls in df[df['relevance'] == 0]['class'].unique(): subset = df[(df['relevance'] == 0) & (df['class'] == cls)] # 根据当前数量选择上/下采样 if len(subset) > target_per_class: adjusted = resample(subset, random_state=SEED, n_samples=target_per_class, replace=False) else: adjusted = resample(subset, random_state=SEED, n_samples=target_per_class, replace=True) majority_groups.append(adjusted) majority_processed = pd.concat(majority_groups) # 合并最终数据集 final_df = pd.concat([minority_processed, majority_processed]) # 验证结果 print("relevance列分布:") print(final_df['relevance'].value_counts()) print("\nclass列分布:") print(final_df['class'].value_counts()) print("\n各(relevance, class)组合分布:") print(final_df.groupby(['relevance', 'class']).size())
验证说明
运行代码后,输出的验证结果会显示:
relevance列两类记录数均为14810左右(若无法整除总数会有细微差异,可根据需求调整目标值)class列三个类别的记录数基本一致- 每个
(relevance, class)组合的记录数相等,完全满足预设的双重约束
内容的提问来源于stack exchange,提问作者mariant
相关产品推荐
相关产品推荐

