如何在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
相关产品推荐
相关产品推荐

