如何通过PyCaret提取混淆矩阵中Tp、Tn、Fp、Fn的数值
PyCaret提取混淆矩阵TP、TN、FP、FN数值方法
当然可以提取具体数值,根本不是只能出可视化图。PyCaret绘制混淆矩阵的逻辑是先计算好底层结构化的统计数值,再基于数值渲染图表,所有图里标注的计数都能通过内置接口直接获取,常用的提取方式有两种:
方法1:直接调用绘图接口返回原始数值
调用plot_model画混淆矩阵的时候,传入return_data=True参数,同时关闭自动弹窗渲染,函数会直接返回Pandas格式的混淆矩阵统计表,不需要手动从图里读数:
# 注意提前完成setup初始化、模型训练流程,替换成你自己训练好的模型对象 cm_table = plot_model( estimator=your_trained_model, plot="confusion_matrix", plot_kwargs={"percent": False}, # 关闭百分比显示,返回原始样本计数 return_data=True, display_format=None # 不弹出可视化绘图窗口 )
返回的cm_table行索引对应真实标签,列索引对应模型预测标签,二分类场景下直接按位置索引就能取出四个值:
- TP(真阳性):
cm_table.iloc[1, 1] - TN(真阴性):
cm_table.iloc[0, 0] - FP(假阳性):
cm_table.iloc[0, 1] - FN(假阴性):
cm_table.iloc[1, 0]
如果是多分类场景,拿到完整混淆矩阵后,对每个类别单独做矩阵切片,就能算出每个类别对应的四个指标值。
方法2:基于预测结果自行统计计算
如果需要更灵活的统计逻辑,可以直接用predict_model拿到验证集/测试集的全量预测结果,手动比对真实标签和预测标签统计数值:
pred_result = predict_model(your_trained_model, data=your_holdout_dataset) # 替换成你数据集里真实标签的列名、对应正负类的标签值即可 TP = ((pred_result["target_column"] == 1) & (pred_result["prediction_label"] == 1)).sum() TN = ((pred_result["target_column"] == 0) & (pred_result["prediction_label"] == 0)).sum() FP = ((pred_result["target_column"] == 0) & (pred_result["prediction_label"] == 1)).sum() FN = ((pred_result["target_column"] == 1) & (pred_result["prediction_label"] == 0)).sum()
注意:如果建模时手动调整过分类概率阈值,用第二种方法计算时要记得同步应用相同的阈值规则,避免统计出来的数值和模型自带评估报告里的结果不一致。
内容的提问来源于stack exchange,提问作者Dmilo6
相关产品推荐
相关产品推荐

