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

如何在Sklearn中获取多分类任务中每个类别的平衡准确率?

多分类任务中获取每个类别的平衡准确率(Sklearn实现)

我正在进行一项多分类任务,使用Python的Sklearn可以获取**准确率(accuracy)与平衡准确率(balanced accuracy)**指标,但二者仅返回单一数值。我希望获取每个类别的平衡准确率——此前用R语言的caret包建模时,其指标报告能输出每个类别的平衡准确率(如下例最后一行),想在Sklearn或相关Python库中实现该功能。


R语言caret包指标示例

执行代码:

confusionMatrix(testSet$classes,testSet$pred_model)

输出结果:

Class: A   Class: B   Class: C  Class: D  Class: E  Class: F
Sensitivity            0.37143   0.23404   0.25490   0.15254   0.30909   0.27692
Specificity            0.85921   0.84528   0.85057   0.83004   0.86381   0.86235
Pos Pred Value         0.25000   0.21154   0.25000   0.17308   0.32692   0.34615
Neg Pred Value         0.91538   0.86154   0.85385   0.80769   0.85385   0.81923
Prevalence             0.11218   0.15064   0.16346   0.18910   0.17628   0.20833
Detection Rate         0.04167   0.03526   0.04167   0.02885   0.05449   0.05769
Detection Prevalence   0.16667   0.16667   0.16667   0.16667   0.16667   0.16667
Balanced Accuracy      0.61532   0.53966   0.55274   0.49129   0.58645   0.56964

Python Sklearn当前实现

执行代码:

from sklearn.metrics import accuracy_score, balanced_accuracy_score

acc = accuracy_score(y_test, y_pred)
print(acc)  # 输出:0.52345

bal_acc = balanced_accuracy_score(y_test, y_pred)
print(bal_acc)  # 输出:0.53657

解决方案:计算每个类别的平衡准确率

平衡准确率的本质是该类别的灵敏度(召回率)与特异度的平均值,公式为:
Balanced Accuracy per class = (Sensitivity + Specificity) / 2

方法1:手动计算单个类别平衡准确率

借助Sklearn的混淆矩阵和召回率函数实现:

from sklearn.metrics import confusion_matrix, recall_score
import numpy as np

# 计算混淆矩阵
cm = confusion_matrix(y_test, y_pred)
n_classes = cm.shape[0]

# 计算每个类别的灵敏度(召回率)
sensitivity = recall_score(y_test, y_pred, average=None)

# 计算每个类别的特异度
specificity = []
for i in range(n_classes):
    # 真阴性:非当前类中被正确分类的样本数
    tn = np.sum(cm) - np.sum(cm[i, :]) - np.sum(cm[:, i]) + cm[i, i]
    # 假阳性:非当前类被错误分类为当前类的样本数
    fp = np.sum(cm[:, i]) - cm[i, i]
    specificity.append(tn / (tn + fp))
specificity = np.array(specificity)

# 计算每个类别的平衡准确率
per_class_bal_acc = (sensitivity + specificity) / 2

# 格式化输出
class_names = ["A", "B", "C", "D", "E", "F"]
for cls, bal_acc in zip(class_names, per_class_bal_acc):
    print(f"Class: {cls} - Balanced Accuracy: {bal_acc:.5f}")

方法2:生成类似caret的完整指标报告

自定义函数生成包含所有分类指标的DataFrame:

from sklearn.metrics import confusion_matrix, recall_score, precision_score
import pandas as pd
import numpy as np

def class_wise_metrics(y_true, y_pred, class_names):
    cm = confusion_matrix(y_true, y_pred)
    n_classes = cm.shape[0]
    
    # 初始化指标列表
    sensitivity = recall_score(y_true, y_pred, average=None)
    precision = precision_score(y_true, y_pred, average=None)
    specificity = []
    neg_pred_value = []
    prevalence = []
    detection_rate = []
    detection_prevalence = []
    
    for i in range(n_classes):
        tp = cm[i, i]
        fn = np.sum(cm[i, :]) - tp
        fp = np.sum(cm[:, i]) - tp
        tn = np.sum(cm) - tp - fn - fp
        
        # 计算各类指标
        spec = tn / (tn + fp) if (tn + fp) != 0 else 0
        npv = tn / (tn + fn) if (tn + fn) != 0 else 0
        prev = (tp + fn) / np.sum(cm)
        det_rate = tp / np.sum(cm)
        det_prev = (tp + fp) / np.sum(cm)
        
        specificity.append(spec)
        neg_pred_value.append(npv)
        prevalence.append(prev)
        detection_rate.append(det_rate)
        detection_prevalence.append(det_prev)
    
    # 计算平衡准确率
    bal_acc = (np.array(sensitivity) + np.array(specificity)) / 2
    
    # 构建结果DataFrame
    metrics_df = pd.DataFrame({
        "Sensitivity": sensitivity,
        "Specificity": specificity,
        "Pos Pred Value": precision,
        "Neg Pred Value": neg_pred_value,
        "Prevalence": prevalence,
        "Detection Rate": detection_rate,
        "Detection Prevalence": detection_prevalence,
        "Balanced Accuracy": bal_acc
    }, index=[f"Class: {cls}" for cls in class_names])
    
    return metrics_df

# 使用示例
class_names = ["A", "B", "C", "D", "E", "F"]
metrics_report = class_wise_metrics(y_test, y_pred, class_names)
print(metrics_report.round(5))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 07:33:12