如何平衡含布尔型标签列的数据集?改写脚本实现泛化
数据集平衡:适配布尔列与通用化实现
我有一个CSV文件,其中包含名为"worked"的列,需要平衡该列值为true和false的行数(使两者数量相等)。我曾编写过针对名为"label"的列、值为二进制0或1的数据集平衡脚本,但不确定如何将其扩展到当前场景,或是实现更通用的泛化处理。原有脚本如下:
# balance the dataset so there are an equal number of 0 and 1 labels import random import pandas as pd INPUT_DATASET = "input_dataset.csv" OUTPUT_DATASET = "output_dataset.csv" LABEL_COL = "label" # load the dataset dataset = pd.read_csv(INPUT_DATASET) # figure out the minimum number of 0s and 1s num_0s = dataset[dataset[LABEL_COL] == 0].shape[0] num_1s = dataset[dataset[LABEL_COL] == 1].shape[0] min_num_rows = min(num_0s, num_1s) print(f"There were {num_0s} 0s and {num_1s} 1s in the dataset - the kept amount is {min_num_rows}.") # randomly select the minumum number of rows for both 0s and 1s chosen_ids = [] for label in (0, 1): ids = dataset[dataset[LABEL_COL] == label].index chosen_ids.extend(random.sample(list(ids), min_num_rows)) # remove the non-chosen ids from the dataset dataset = dataset.drop(dataset.index[list(set(range(dataset.shape[0])) - set(chosen_ids))]) # save the dataset dataset.to_csv(OUTPUT_DATASET, index=False)
一、快速适配布尔列的修改
只需要调整原有脚本中的几个关键部分,就能适配worked列的true/false平衡需求:
# 平衡数据集,使worked列的true和false行数相等 import random import pandas as pd INPUT_DATASET = "input_dataset.csv" OUTPUT_DATASET = "output_dataset.csv" LABEL_COL = "worked" # 替换为目标列名 # 加载数据集 dataset = pd.read_csv(INPUT_DATASET) # 统计true和false的数量,取最小值作为保留行数 num_true = dataset[dataset[LABEL_COL] == True].shape[0] num_false = dataset[dataset[LABEL_COL] == False].shape[0] min_num_rows = min(num_true, num_false) print(f"原数据集中有{num_true}条true记录和{num_false}条false记录,将各保留{min_num_rows}条。") # 随机抽取对应数量的行索引 chosen_ids = [] for label in (True, False): ids = dataset[dataset[LABEL_COL] == label].index chosen_ids.extend(random.sample(list(ids), min_num_rows)) # 筛选出选中的行(简化原有删除逻辑) dataset = dataset.loc[chosen_ids] # 保存平衡后的数据集 dataset.to_csv(OUTPUT_DATASET, index=False)
二、通用化数据集平衡实现
如果需要处理更多类型的分类列(比如字符串标签、多分类场景),可以封装成通用函数,自动适配任意分类列:
import random import pandas as pd def balance_dataset(input_path, output_path, label_col): # 加载数据集 dataset = pd.read_csv(input_path) # 获取标签列的所有唯一类别 unique_labels = dataset[label_col].unique() # 统计每个类别的行数,找到最小行数 label_counts = dataset[label_col].value_counts() min_count = label_counts.min() print(f"各标签数量:{label_counts.to_dict()},将统一保留{min_count}条。") # 对每个类别随机抽取min_count条记录 balanced_dfs = [] for label in unique_labels: subset = dataset[dataset[label_col] == label] # 随机抽样,random_state保证可复现(可选) sampled_subset = subset.sample(n=min_count, random_state=random.randint(0, 1000)) balanced_dfs.append(sampled_subset) # 合并所有抽样后的子集 balanced_dataset = pd.concat(balanced_dfs, ignore_index=True) # 打乱数据集顺序(可选,避免类别集中排列) balanced_dataset = balanced_dataset.sample(frac=1, random_state=random.randint(0, 1000)).reset_index(drop=True) # 保存结果 balanced_dataset.to_csv(output_path, index=False) print(f"平衡后的数据集已保存至{output_path}") # 使用示例 INPUT_DATASET = "input_dataset.csv" OUTPUT_DATASET = "output_dataset.csv" LABEL_COL = "worked" # 可替换为任意分类列名 balance_dataset(INPUT_DATASET, OUTPUT_DATASET, LABEL_COL)
通用函数优势
- 自动适配任意分类类型:布尔值、整数、字符串标签都能处理
- 支持多分类场景:超过两个类别时,统一将每个类别抽样到最少类别的数量
- 代码更简洁,复用性强,无需每次修改判断逻辑
- 可选打乱数据集顺序,避免后续模型训练时的类别偏差
内容的提问来源于stack exchange,提问作者Pro Q
相关产品推荐
相关产品推荐

