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

skplt.plot_roc触发ValueError:输入变量样本数不一致问题求助

问题描述

运行Python机器学习代码时,调用skplt.metrics.plot_roc(y_test,y_probas)触发ValueError,错误信息如下:

ValueError: Found input variables with inconsistent numbers of samples: [125720, 132006]

已尝试的操作:

  • 将LogisticRegression的max_iter超参数从100调整为500;
  • 打印验证x_train、x_test、y_train、y_test样本数均一致;
  • 打印得y_test形状为(3143, ),y_probas形状为(3143,42)。

相关代码:

# 模型构建与预测函数
def Model(model, X, y):
    # 划分训练测试集
    print(X.shape)
    print(y.shape)
    x_train, x_test, y_train, y_test = train_test_split(X, y, test_size=0.25, random_state=30)
    print(x_train.shape)
    print(y_train.shape)
    print(x_test.shape)
    print(y_test.shape)
    # 构建管道模型:CountVectorizer + TfidfTransformer + 分类器
    pipeline_model = Pipeline([('vect', CountVectorizer()),
                              ('tfidf', TfidfTransformer()),
                              ('clf', model)])
    pipeline_model.fit(x_train, y_train)
    
    y_pred = pipeline_model.predict(x_test)
    y_probas = pipeline_model.predict_proba(x_test)
    print(y_test.shape)
    print(y_probas.shape)
    # 绘制ROC曲线
    skplt.metrics.plot_roc(y_test,y_probas,figsize=(12,8),title_fontsize=12,text_fontsize=16)
    plt.show()
    # 绘制精确率-召回率曲线
    skplt.metrics.plot_precision_recall(y_test,y_probas,figsize=(12,8),title_fontsize=12,text_fontsize=16)
    plt.show()
    # 输出评估指标
    print("混淆矩阵:\n",confusion_matrix(y_test,y_pred))
    print("分类报告:\n",classification_report(y_test, y_pred))
    print('准确率:', pipeline_model.score(x_test, y_test)*100)
    print("训练集得分:\n",pipeline_model.score(x_train,y_train)*100)

# 逻辑回归模型
from sklearn.linear_model import LogisticRegression
model = LogisticRegression(max_iter=500)
Model(model, X, y)
问题排查与解决方案

核心原因

错误提示的样本数[125720,132006]与打印的(3143, )不符,说明plot_roc内部处理标签时出现维度错误。skplot的plot_roc在多分类场景下,默认要求标签为one-hot编码格式,而你的y_test是一维类别索引数组,当类别数为42时,内部转换逻辑可能计算出错,导致样本数匹配失败。

解决办法

  1. 将y_test转换为one-hot编码
    使用sklearn的编码器将一维标签转为one-hot格式后传入:

    from sklearn.preprocessing import OneHotEncoder
    import numpy as np
    
    # 转换y_test为one-hot编码
    encoder = OneHotEncoder(sparse_output=False)
    y_test_onehot = encoder.fit_transform(y_test.reshape(-1, 1))
    
    # 传入one-hot后的标签绘制ROC曲线
    skplt.metrics.plot_roc(y_test_onehot, y_probas, figsize=(12,8), title_fontsize=12, text_fontsize=16)
    
  2. 升级scikit-plot版本
    旧版本的scikit-plot在多分类ROC处理上存在bug,升级到最新版本可修复:

    pip install scikit-plot --upgrade
    
  3. 验证维度一致性
    在调用plot_roc前添加断言,确认输入维度无问题:

    assert y_test.shape[0] == y_probas.shape[0], "输入样本数不匹配"
    

    若断言通过,说明问题出在skplot内部逻辑,优先尝试前两种方法。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 23:34:55