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

多分类ROC曲线绘制报错:无法处理多分类与多标签混合目标

多分类任务ROC曲线绘制问题解决

错误原因分析

你遇到的ValueError主要来自三个核心问题:

  1. 二分类函数硬编码:calculate_tpr_fpr是为二分类设计的,直接取混淆矩阵的cm[0,0]等索引,但多分类(3类)的混淆矩阵是3×3,逻辑完全不匹配。
  2. 语法逻辑错误:get_all_roc_coordinates里的y_pred = y_proba = threshold是严重错误,会把y_proba(样本×类别的概率矩阵)覆盖成单个阈值,导致y_pred格式和y_test(多分类标签)不兼容。
  3. 多分类ROC策略缺失:多分类任务不能直接套用二分类的ROC计算逻辑,必须采用**One-vs-Rest(OvR)或One-vs-One(OvO)**策略,将每个类别单独转化为二分类问题处理。

解决方案:基于OvR策略的ROC曲线绘制

方法1:使用sklearn内置函数(推荐)

sklearn提供了直接支持多分类ROC的工具,无需手动实现计算逻辑,更可靠高效。

完整代码示例:

import numpy as np
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier
from sklearn.metrics import roc_curve, auc
from sklearn.preprocessing import label_binarize

# 拆分数据集并训练模型
X_train, X_test, y_train, y_test = train_test_split(X, Y, random_state=0)
dtree_model = DecisionTreeClassifier(max_depth=2).fit(X_train, y_train)
y_proba = dtree_model.predict_proba(X_test)

# 将多分类标签二值化(OvR策略:每个类别单独作为正类,其余为负类)
n_classes = 3
y_test_binarized = label_binarize(y_test, classes=[0, 1, 2])

# 计算每个类别的FPR、TPR和AUC值
fpr = dict()
tpr = dict()
roc_auc = dict()
for i in range(n_classes):
    fpr[i], tpr[i], _ = roc_curve(y_test_binarized[:, i], y_proba[:, i])
    roc_auc[i] = auc(fpr[i], tpr[i])

# 绘制ROC曲线
plt.figure(figsize=(8,6))
colors = ['#1f77b4', '#ff7f0e', '#2ca02c']
for i, color in zip(range(n_classes), colors):
    plt.plot(fpr[i], tpr[i], color=color, lw=2,
             label=f'类别{i}的ROC曲线 (AUC = {roc_auc[i]:.2f})')

# 添加随机猜测的参考线
plt.plot([0, 1], [0, 1], 'k--', lw=2)
plt.xlim([0.0, 1.0])
plt.ylim([0.0, 1.05])
plt.xlabel('假阳性率(FPR)')
plt.ylabel('真阳性率(TPR)')
plt.title('多分类任务ROC曲线(One-vs-Rest)')
plt.legend(loc="lower right")
plt.show()

方法2:手动实现OvR逻辑(适合学习原理)

如果你需要手动实现计算逻辑,可按以下方式修正函数:

修正后的二分类TPR/FPR计算函数(支持指定正类别)

from sklearn.metrics import confusion_matrix

def calculate_tpr_fpr(y_test, y_pred_bin):
    cm = confusion_matrix(y_test, y_pred_bin)
    TN = cm[0, 0] if cm.shape[0] >=2 else 0
    FP = cm[0, 1] if cm.shape[1] >=2 else 0
    FN = cm[1, 0] if cm.shape[0] >=2 else 0
    TP = cm[1, 1] if cm.shape[1] >=2 else 0
    
    tpr = TP / (TP + FN) if (TP + FN) != 0 else 0
    fpr = 1 - TN / (TN + FP) if (TN + FP) != 0 else 0
    return tpr, fpr

单个类别的ROC坐标计算函数

def get_roc_coordinates_single_class(y_test, y_proba_class, positive_class):
    tpr_list = [0]
    fpr_list = [0]
    # 获取所有唯一概率值作为阈值,倒序排列
    thresholds = np.unique(y_proba_class)
    thresholds = np.sort(thresholds)[::-1]
    
    for threshold in thresholds:
        # 将当前类别转化为二分类标签
        y_test_bin = (y_test == positive_class).astype(int)
        y_pred_bin = (y_proba_class >= threshold).astype(int)
        tpr, fpr = calculate_tpr_fpr(y_test_bin, y_pred_bin)
        tpr_list.append(tpr)
        fpr_list.append(fpr)
    
    return tpr_list, fpr_list

调用示例

# 计算类别0的ROC坐标
tpr_0, fpr_0 = get_roc_coordinates_single_class(y_test, y_proba[:,0], positive_class=0)
# 计算类别1的ROC坐标
tpr_1, fpr_1 = get_roc_coordinates_single_class(y_test, y_proba[:,1], positive_class=1)
# 计算类别2的ROC坐标
tpr_2, fpr_2 = get_roc_coordinates_single_class(y_test, y_proba[:,2], positive_class=2)

# 绘制曲线(代码同方法1的绘图部分,替换对应数据即可)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 02:48:24