多分类ROC曲线绘制报错:无法处理多分类与多标签混合目标
多分类任务ROC曲线绘制问题解决
错误原因分析
你遇到的ValueError主要来自三个核心问题:
- 二分类函数硬编码:
calculate_tpr_fpr是为二分类设计的,直接取混淆矩阵的cm[0,0]等索引,但多分类(3类)的混淆矩阵是3×3,逻辑完全不匹配。 - 语法逻辑错误:
get_all_roc_coordinates里的y_pred = y_proba = threshold是严重错误,会把y_proba(样本×类别的概率矩阵)覆盖成单个阈值,导致y_pred格式和y_test(多分类标签)不兼容。 - 多分类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
相关产品推荐
相关产品推荐

