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

如何用Python/Scikit-learn计算多分类各类别的Sensitivity、Specificity等指标

多分类场景下按类别计算Sensitivity、Specificity和Positive Predictivity的Python方案

Scikit-learn没有直接提供多分类下每个类别这三个指标的开箱即用API,但可以基于confusion_matrix的结果手动推导,或结合内置函数快速实现,以下是针对4分类场景的具体方案:

核心逻辑

对多分类中的每个类别,将其视为二分类问题(当前类别为正类,其余所有类别合并为负类),再通过混淆矩阵计算基础指标后推导目标值:

  • Sensitivity(召回率):TP / (TP + FN),对应sklearn的recall_score(按类别计算)
  • Specificity:TN / (TN + FP),需通过混淆矩阵全局统计推导
  • Positive Predictivity(精确率):TP / (TP + FP),对应sklearn的precision_score(按类别计算)

完整实现代码

1. 生成混淆矩阵

先通过confusion_matrix获取真实标签与预测标签的混淆矩阵:

from sklearn.metrics import confusion_matrix
import numpy as np

# 示例4分类数据:真实标签和预测标签(类别0-3)
y_true = [0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3]
y_pred = [0, 1, 1, 3, 0, 2, 2, 3, 1, 1, 2, 3]

# 生成混淆矩阵
cm = confusion_matrix(y_true, y_pred)

2. 自定义函数计算全类别指标

遍历每个类别,基于混淆矩阵计算TP、TN、FP、FN,再推导目标指标:

def calculate_multiclass_metrics(confusion_matrix):
    num_classes = confusion_matrix.shape[0]
    class_metrics = {}
    
    for cls in range(num_classes):
        # 提取当前类别的基础指标
        tp = confusion_matrix[cls, cls]
        fn = confusion_matrix[cls, :].sum() - tp
        fp = confusion_matrix[:, cls].sum() - tp
        tn = confusion_matrix.sum() - (tp + fn + fp)
        
        # 计算目标指标(避免除以0错误)
        sensitivity = tp / (tp + fn) if (tp + fn) != 0 else 0
        specificity = tn / (tn + fp) if (tn + fp) != 0 else 0
        positive_predictivity = tp / (tp + fp) if (tp + fp) != 0 else 0
        
        class_metrics[f"类别{cls}"] = {
            "Sensitivity": round(sensitivity, 4),
            "Specificity": round(specificity, 4),
            "Positive Predictivity": round(positive_predictivity, 4)
        }
    
    return class_metrics

# 计算并打印结果
metrics = calculate_multiclass_metrics(cm)
for cls, vals in metrics.items():
    print(f"{cls}:")
    for name, val in vals.items():
        print(f"  {name}: {val}")
    print()

3. 用Scikit-learn内置函数验证部分指标

Sensitivity和Positive Predictivity可直接用内置函数快速计算,验证自定义结果:

from sklearn.metrics import precision_score, recall_score

# 按类别计算精确率(Positive Predictivity)
precision = precision_score(y_true, y_pred, average=None)
# 按类别计算召回率(Sensitivity)
recall = recall_score(y_true, y_pred, average=None)

print("内置函数验证结果:")
for cls in range(4):
    print(f"类别{cls}:")
    print(f"  Recall(Sensitivity): {round(recall[cls], 4)}")
    print(f"  Precision(Positive Predictivity): {round(precision[cls], 4)}")
    print()

注意事项

  • 当某类别无真实样本(TP+FN=0)或无预测样本(TP+FP=0)时,指标默认设为0,避免除以0异常
  • 该方案支持任意多分类场景,只需保证混淆矩阵的维度与类别数量匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 17:35:23