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

使用scikit-learn的roc_auc_score时遇多分类及维度错误求助

解决scikit-learn v1.3 roc_auc_score的多分类与维度错误

错误原因分析

  1. 初始多分类错误:你的y_label包含3个类别(0、1、2),scikit-learn自动判定为多分类任务,但你未指定multi_class参数,因此触发ValueError。
  2. 维度错误:添加multi_class后,函数要求y_score(即你的pred_proba)为2D数组(形状为(n_samples, n_classes),每个样本对应所有类别的概率/得分),但你的pred_proba是1D数组,导致AxisError。

解决方案

根据你的任务需求选择以下任一方案:

方案1:转化为二分类任务计算单类别ROC AUC

如果只关注某个特定类别的性能,可以将多分类任务转化为二分类(目标类别为正类,其余为负类):

# 示例:将类别2设为正类,其他类别设为负类
y_binary = (y_label == 2).astype(int)
# 假设pred_proba是类别2的得分/概率值
roc_auc = roc_auc_score(y_binary, pred_proba)
print("ROC (%)", 100 * roc_auc)

方案2:使用多分类ROC AUC计算(需2D得分数组)

如果需要评估所有类别的整体性能,需确保pred_proba是每个样本对应所有类别的概率/得分(2D数组):

  1. 重新获取模型输出:如果使用scikit-learn分类器,调用predict_proba()而非predict(),后者仅输出类别标签,前者输出所有类别的概率:
# 假设你的模型对象为clf,测试集为X_test
pred_proba = clf.predict_proba(X_test)  # 形状为(n_samples, 3)
  1. 调用roc_auc_score并指定多分类模式:
roc_auc = roc_auc_score(y_label, pred_proba, multi_class='ovr')  # 或使用'ovo'
print("ROC (%)", 100 * roc_auc)

关键说明

  • multi_class='ovr':计算每个类别相对于其他类别的ROC AUC,取平均值;
  • multi_class='ovo':计算每对类别之间的ROC AUC,取平均值;
  • 1D的pred_proba仅适用于二分类任务,多分类任务必须提供2D的得分/概率数组。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 14:32:31