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

如何生成仅含置信度超阈值预测的混淆矩阵?Sklearn支持吗?

解决方法:过滤高置信度样本后生成混淆矩阵

好问题!我之前也碰到过类似的需求——sklearn的confusion_matrix确实没有直接提供置信度过滤的参数,但手动实现起来非常简单,咱们一步步来:

一、核心思路

我们需要先筛选出预测最大概率值超过阈值的样本,然后用这些样本的真实标签和预测标签生成混淆矩阵。关键是要生成一个布尔掩码,同时过滤真实标签数组和预测标签数组,确保两者的样本一一对应。

二、具体实现步骤(附代码)

假设你已经有了独热编码的真实标签y_TEST和模型输出的预测概率y_pred,可以按以下步骤操作:

1. 导入必要工具

import numpy as np
from sklearn.metrics import confusion_matrix

2. 将独热编码转成类别索引

confusion_matrix需要的是类别索引(比如0/1这类离散值),而非独热向量,所以先做转换:

# 真实标签的类别索引
y_true = y_TEST.argmax(axis=1)
# 基于最大概率得到的预测类别索引
y_pred_classes = y_pred.argmax(axis=1)

3. 生成高置信度样本的掩码

计算每个样本的最大预测概率,再筛选出符合阈值要求的样本:

# 计算每个样本的最大置信度
max_confidences = np.max(y_pred, axis=1)
# 设置阈值,生成布尔掩码(True表示该样本符合高置信度要求)
confidence_threshold = 0.9
mask = max_confidences > confidence_threshold

4. 过滤标签并生成混淆矩阵

用掩码同时过滤真实标签和预测标签,再传入confusion_matrix即可:

# 过滤后的真实标签与预测标签
y_true_filtered = y_true[mask]
y_pred_classes_filtered = y_pred_classes[mask]

# 生成过滤后的混淆矩阵(支持normalize参数,用法和原生函数一致)
filtered_confusion_matrix = confusion_matrix(
    y_true_filtered, 
    y_pred_classes_filtered, 
    normalize='pred'  # 可按需替换为'true'/'all'
)

print(filtered_confusion_matrix)

三、可选:封装成复用函数

如果经常需要使用这个功能,可以封装成一个函数,方便后续调用:

def confidence_filtered_confusion_matrix(y_test_onehot, y_pred_proba, threshold=0.9, normalize=None):
    y_true = y_test_onehot.argmax(axis=1)
    y_pred_classes = y_pred_proba.argmax(axis=1)
    max_conf = np.max(y_pred_proba, axis=1)
    mask = max_conf > threshold
    return confusion_matrix(y_true[mask], y_pred_classes[mask], normalize=normalize)

# 调用示例
my_confusion_matrix = confidence_filtered_confusion_matrix(y_TEST, y_pred, threshold=0.9, normalize='pred')

这样就完全实现了你想要的「仅保留最大预测值大于阈值的样本生成混淆矩阵」的效果啦!

内容的提问来源于stack exchange,提问作者TommyF

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 19:52:39