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

多分类(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 01:01:01