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

调用roc_curve绘制ROC曲线报multiclass format is not supported错误

报错触发原因

scikit-learn库内置的roc_curve函数仅原生支持二分类任务的ROC指标计算,出现multiclass format is not supported报错的核心原因是入参不符合二分类要求:要么传入的y_test是多分类标签格式,要么传入的out_pred_prob是多分类模型输出的形状为(样本数, 类别数)的二维概率矩阵,函数无法直接处理多分类输入。

修复方案

根据实际任务类型选择对应修复方式:

  • 任务为二分类,入参格式传错
    二分类场景下roc_curve要求传入的预测值是一维的正类概率数组,形状为(样本数,),不能传入同时包含正负两类概率的二维数组。如果out_pred_prob是模型输出的两列概率矩阵,取索引为1的列(对应正类的预测概率)传入即可:

    # 仅修改roc_curve入参,取正类(类别索引为1)的预测概率
    fpr, tpr, thresholds = roc_curve(y_test, out_pred_prob[:, 1])
    # 后续绘图逻辑保持原有写法不变
    plt.plot(fpr, tpr, label='ROC curve')
    plt.plot([0, 1], [0, 1], 'k--', label='Random guess')
    _ = plt.xlabel('False Positive Rate')
    _ = plt.ylabel('True Positive Rate')
    _ = plt.title('ROC Curve')
    _ = plt.xlim([-0.02, 1])
    _ = plt.ylim([0, 1.02])
    _ = plt.legend(loc="lower right")
    
  • 任务为多分类,需适配多分类ROC计算逻辑
    多分类场景无法直接调用原生roc_curve,需要先将多分类标签做二值化转换,再通过一对多(One-vs-Rest)的逻辑分别计算每个类别的ROC指标,也可以计算宏/微平均的聚合ROC曲线。参考实现代码:

    from sklearn.preprocessing import label_binarize
    from sklearn.metrics import roc_curve, auc
    import numpy as np
    
    # 按实际类别数量修改n_classes取值
    n_classes = 3
    # 对测试集多分类标签做二值化处理
    y_test_bin = label_binarize(y_test, classes=[*range(n_classes)])
    
    fpr = dict()
    tpr = dict()
    roc_auc = dict()
    # 逐类别计算ROC指标
    for i in range(n_classes):
        fpr[i], tpr[i], _ = roc_curve(y_test_bin[:, i], out_pred_prob[:, i])
        roc_auc[i] = auc(fpr[i], tpr[i])
    
    # 计算微平均聚合ROC
    fpr["micro"], tpr["micro"], _ = roc_curve(y_test_bin.ravel(), out_pred_prob.ravel())
    roc_auc["micro"] = auc(fpr["micro"], tpr["micro"])
    
    # 绘图
    plt.figure()
    # 逐类绘制ROC曲线
    for i in range(n_classes):
        plt.plot(fpr[i], tpr[i], label=f'ROC curve of class {i} (area = {roc_auc[i]:0.2f})')
    # 绘制聚合ROC与基准线
    plt.plot(fpr["micro"], tpr["micro"], label=f'Micro-average ROC (area = {roc_auc["micro"]:0.2f})', linestyle='--')
    plt.plot([0, 1], [0, 1], 'k--', label='Random guess')
    _ = plt.xlabel('False Positive Rate')
    _ = plt.ylabel('True Positive Rate')
    _ = plt.title('Multi-class ROC Curve')
    _ = plt.xlim([-0.02, 1])
    _ = plt.ylim([0, 1.02])
    _ = plt.legend(loc="lower right")
    plt.show()
    

    如果只需要单条聚合ROC曲线,直接使用微平均的fpr、tpr结果绘制即可,不需要逐类输出曲线。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 22:15:45