多分类(7标签)CNN模型绘制ROC曲线遇格式错误求解决方案
解决多分类ROC曲线绘制的ValueError问题
问题原因
sklearn.metrics.roc_curve默认仅支持二分类任务,而你的数据是7分类的one-hot编码格式(y_test和y_pred形状为(样本数, 7)),直接传入会触发multilabel-indicator format is not supported错误。
解决方案:按类别逐一绘制ROC曲线
针对多分类场景,我们需要将每个类别视为独立的二分类任务(属于该类/不属于该类),逐一计算并绘制ROC曲线。
步骤1:导入依赖库
import matplotlib.pyplot as plt from sklearn import metrics
步骤2:逐类别计算ROC指标
假设y_test和y_pred均为one-hot编码格式(每列对应一个类别的标签/预测概率):
n_classes = 7 # 你的标签总数 fpr = dict() tpr = dict() roc_auc = dict() # 遍历每个类别,计算对应的FPR、TPR和AUC for i in range(n_classes): fpr[i], tpr[i], _ = metrics.roc_curve(y_test[:, i], y_pred[:, i]) roc_auc[i] = metrics.auc(fpr[i], tpr[i])
如果y_test是类别索引格式(形状为(样本数,)的整数数组),需要先将其转换为one-hot编码:
from sklearn.preprocessing import label_binarize # 将类别索引二值化为one-hot格式 y_test_binarized = label_binarize(y_test, classes=[0, 1, 2, 3, 4, 5, 6]) n_classes = y_test_binarized.shape[1] # 再执行上述逐类别计算逻辑 for i in range(n_classes): fpr[i], tpr[i], _ = metrics.roc_curve(y_test_binarized[:, i], y_pred[:, i]) roc_auc[i] = metrics.auc(fpr[i], tpr[i])
步骤3:绘制多类别ROC曲线
plt.figure(figsize=(8, 6)) # 定义不同类别的颜色 colors = ['blue', 'red', 'green', 'yellow', 'cyan', 'magenta', 'black'] # 绘制每个类别的ROC曲线 for i, color in zip(range(n_classes), colors): plt.plot(fpr[i], tpr[i], color=color, lw=2, label=f'类别{i} (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曲线') plt.legend(loc='lower right') plt.show()
可选:计算宏平均/微平均ROC
如果需要整体评估模型性能,可以计算宏平均或微平均的ROC指标:
# 微平均ROC(将所有类别预测合并为二分类问题) fpr["micro"], tpr["micro"], _ = metrics.roc_curve(y_test.ravel(), y_pred.ravel()) roc_auc["micro"] = metrics.auc(fpr["micro"], tpr["micro"]) # 宏平均AUC(对每个类别AUC取算术平均) roc_auc["macro"] = metrics.roc_auc_score(y_test, y_pred, average="macro") # 绘制微平均ROC曲线 plt.figure() plt.plot(fpr["micro"], tpr["micro"], label=f'微平均ROC (AUC = {roc_auc["micro"]:.2f})', color='deeppink', linestyle=':', linewidth=4) plt.plot([0, 1], [0, 1], 'k--', lw=2) plt.xlabel('假阳性率(FPR)') plt.ylabel('真阳性率(TPR)') plt.title('微平均ROC曲线') plt.legend(loc='lower right') plt.show()
内容的提问来源于stack exchange,提问作者PKrange
相关产品推荐
相关产品推荐

