如何构建覆盖所有类别不平衡组合的ROC/F1分析数据集
固定总样本量下全合法混淆矩阵组合高效生成方案
实现思路
- 核心约束:所有合法的混淆矩阵取值必须满足
TP、TN、FP、FN均为非负整数,且TP + TN + FP + FN = 100。原实现仅固定TP=0、TN=0,仅覆盖了101种极端场景,完全无法支撑全场景可视化需求。 - 降维提效:四个变量存在线性约束,不需要4层全遍历,仅需遍历TP、FP、FN三个变量的0~100取值范围,TN通过总样本量公式直接计算,自动过滤TN为负的非法组合,比逐行循环的生成效率高2个数量级以上。
- 前置校验:生成四值后同步预计算所有可视化需要的衍生指标,提前处理分母为0的边界情况(比如实际无正样本、预测无正样本的极端场景),从根源避免可视化时出现除零报错、指标超出[0,1]范围的逻辑错误。
- 适配交互:生成的数据集体量轻量,所有指标预计算完成,可直接绑定交互滑块组件的参数映射,不需要实时计算,交互无卡顿。
实现代码
import numpy as np import pandas as pd # 固定总样本量参数,可按需调整 TOTAL_OBS = 100 # 生成三个独立变量的合法整数取值范围 tp = np.arange(TOTAL_OBS + 1) fp = np.arange(TOTAL_OBS + 1) fn = np.arange(TOTAL_OBS + 1) # 向量化生成所有取值组合,避免低效率循环 TP, FP, FN = np.meshgrid(tp, fp, fn, indexing='ij') TP = TP.ravel() FP = FP.ravel() FN = FN.ravel() # 按总样本约束计算TN,过滤所有非法负值组合 TN = TOTAL_OBS - TP - FP - FN valid_mask = TN >= 0 TP, TN, FP, FN = [arr[valid_mask] for arr in [TP, TN, FP, FN]] # 预计算ROC、F1需要的所有衍生指标,提前处理分母为0的边界 with np.errstate(divide='ignore', invalid='ignore'): tpr = np.where(TP + FN == 0, 0, TP / (TP + FN)) # 真阳性率,同召回率 fpr = np.where(FP + TN == 0, 0, FP / (FP + TN)) # 假阳性率 precision = np.where(TP + FP == 0, 0, TP / (TP + FP)) # 精确率 f1 = np.where(precision + tpr == 0, 0, 2 * precision * tpr / (precision + tpr)) # F1分数 # 组装最终数据集 cm_dataset = pd.DataFrame({ 'TP': TP, 'TN': TN, 'FP': FP, 'FN': FN, 'TPR': tpr, 'FPR': fpr, 'precision': precision, 'recall': tpr, 'f1_score': f1 })
方案说明
- 最终生成的数据集共包含176851条合法组合,完全覆盖总样本量为100时所有可能的混淆矩阵场景,无遗漏、无非法值。
- 若需要缩减数据集体量提升交互流畅度,可调整
np.arange的步长参数做等间隔采样,不需要修改核心逻辑。 - 所有字段取值范围均符合指标定义,可直接用于ROC曲线、F1-召回率/精确率曲线的动态渲染,不需要额外做数据清洗。
内容的提问来源于stack exchange,提问作者Chris
相关产品推荐
相关产品推荐

