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

如何在Python中排查多CSV文件类别不平衡及解决代码报错

解决CSV文件合并统计类别不平衡问题

问题背景

需要排查10个CSV训练文件中的类别不平衡问题,计划将所有文件的pXC50列(即y_train)合并为纵向长列,通过value_counts()统计类别分布,但合并时触发错误:

first argument must be an iterable of pandas objects, you passed an object of type "Series"

原错误代码:

import pandas as pd

datasets = ['CHEMBL4794', 'CHEMBL4805', 'CHEMBL4822', 'CHEMBL1293228', 'CHEMBL1741171',
            'CHEMBL1907607', 'CHEMBL1907608', 'CHEMBL1907610', 'CHEMBL2093869', 
            'CHEMBL2094108']

all_y_train = []

def run_check(dataset): 
    df_train = pd.read_csv(f'C:\\Users\\AMahmud\\Classification_Analysis\\input\\base_processed\\{dataset}_train.csv')

    x_train = df_train.drop(columns = ['molecule_id','pXC50'])
    y_train = (df_train.pXC50)
    all_y_train = pd.concat(y_train)
    

def check_balance():
    for dataset in datasets:
        run_check(dataset)
        
check_balance()

错误原因与修复

错误分析

pd.concat()要求传入可迭代的pandas对象集合(如装多个Series/DataFrame的列表),但你直接传入了单个Series;同时函数内的all_y_train = pd.concat(y_train)会覆盖全局的列表变量,无法累积多个数据集的y_train。

修正后的代码

import pandas as pd

datasets = ['CHEMBL4794', 'CHEMBL4805', 'CHEMBL4822', 'CHEMBL1293228', 'CHEMBL1741171',
            'CHEMBL1907607', 'CHEMBL1907608', 'CHEMBL1907610', 'CHEMBL2093869', 
            'CHEMBL2094108']

all_y_train = []

def run_check(dataset): 
    global all_y_train  # 声明使用全局列表变量
    df_train = pd.read_csv(f'C:\\Users\\AMahmud\\Classification_Analysis\\input\\base_processed\\{dataset}_train.csv')
    y_train = df_train['pXC50']  # 规范的列索引写法
    all_y_train.append(y_train)  # 将当前数据集的y_train加入全局列表

def check_balance():
    for dataset in datasets:
        run_check(dataset)
    # 所有数据集处理完成后,一次性合并列表中的所有Series
    combined_y = pd.concat(all_y_train, ignore_index=True)
    # 输出类别计数与占比
    print("=== 整体训练集类别分布 ===")
    print(combined_y.value_counts())
    print("\n类别占比:")
    print(combined_y.value_counts(normalize=True))
        
check_balance()

更高效排查类别不平衡的方法

1. 同时统计单个数据集与整体的分布

不要只看合并后的整体结果,逐个数据集排查能快速定位不平衡源头,修改run_check函数即可:

def run_check(dataset): 
    global all_y_train
    df_train = pd.read_csv(f'C:\\Users\\AMahmud\\Classification_Analysis\\input\\base_processed\\{dataset}_train.csv')
    y_train = df_train['pXC50']
    all_y_train.append(y_train)
    # 打印当前数据集的分布细节
    print(f"\n=== {dataset} 类别分布 ===")
    print(y_train.value_counts())
    print("类别占比:")
    print(y_train.value_counts(normalize=True))

2. 可视化类别分布

用柱状图直观展示分布情况,比纯数字更易发现问题:

import seaborn as sns
import matplotlib.pyplot as plt

# 绘制整体分布
combined_y = pd.concat(all_y_train, ignore_index=True)
plt.figure(figsize=(8,5))
sns.countplot(x=combined_y)
plt.title("所有训练集合并后的类别分布")
plt.xlabel("pXC50 类别")
plt.ylabel("样本数量")
plt.show()

# 绘制每个数据集的子图对比
fig, axes = plt.subplots(5, 2, figsize=(16,20))
axes = axes.flatten()
for idx, dataset in enumerate(datasets):
    df_train = pd.read_csv(f'C:\\Users\\AMahmud\\Classification_Analysis\\input\\base_processed\\{dataset}_train.csv')
    y_train = df_train['pXC50']
    sns.countplot(x=y_train, ax=axes[idx])
    axes[idx].set_title(f"{dataset} 类别分布")
    axes[idx].tick_params(axis='x', rotation=45)  # 防止类别名重叠
plt.tight_layout()
plt.show()

3. 计算类别不平衡权重

通过sklearn的工具计算类别权重,量化不平衡程度:

from sklearn.utils.class_weight import compute_class_weight

combined_y = pd.concat(all_y_train, ignore_index=True)
# 计算balanced模式下的类别权重(权重越高说明样本越少)
class_weights = compute_class_weight('balanced', classes=combined_y.unique(), y=combined_y)
weight_dict = dict(zip(combined_y.unique(), class_weights))
print("\n类别不平衡权重:")
print(weight_dict)

内容的提问来源于stack exchange,提问作者Adnan Mahmud

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 10:24:50