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

如何平衡含布尔型标签列的数据集?改写脚本实现泛化

数据集平衡:适配布尔列与通用化实现

我有一个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 22:39:25