sklearn绘制RocCurve报错str与int无法比较的问题求解
问题根因
这个TypeError本质是参与大小比较的值类型不统一:roc_curve内部会遍历阈值对真实标签、预测得分做布尔判断,只要传入的y_true真实标签数组里混入了字符串类型值就会触发报错,和模型、预处理逻辑的算法本身无关,基本都是数据拆分、列筛选环节的疏漏导致的。
排查&修复步骤
- 第一步先核查
roc_curve入参的真实标签:直接执行print(y_test.dtype, set(y_test)),如果输出类型不是int/float、集合里出现字符串值,就说明标签被污染了。最高发的错误:构造特征矩阵X时没有剔除目标列
PRZYPADKOWE_CZY_CELOWE,甚至把用户ID、备注这类纯文本非特征列放进了特征集,导致列筛选、预处理环节特征和标签错位,最终y_test里混入字符串。 - 第二步核查列选择器的筛选范围:不要直接用默认规则全量扫列,显式指定排除非特征字段,避免把不该进模型的列捞进去:
from sklearn.compose import make_column_selector import numpy as np # 显式筛选数值列,排除非特征字段 num_selector = make_column_selector(dtype_include=np.number, exclude=['PRZYPADKOWE_CZY_CELOWE']) # 显式筛选类别列,排除非特征字段 cat_selector = make_column_selector(dtype_include=object, exclude=['PRZYPADKOWE_CZY_CELOWE']) - 第三步核查
decision_function的输出:执行print(y_score.shape, y_score.dtype),正常应该是长度和测试集样本数一致的一维浮点数组,如果形状和测试集样本数对不上,说明前面Pipeline输入的特征维度错误,一定混了多余列。
可直接运行的正确代码示例
import pandas as pd import numpy as np from sklearn.model_selection import train_test_split from sklearn.compose import ColumnTransformer, make_column_selector from sklearn.preprocessing import OneHotEncoder, StandardScaler from sklearn.linear_model import LogisticRegression from sklearn.pipeline import make_pipeline from sklearn.metrics import roc_curve, RocCurveDisplay # 1. 显式拆分特征和标签,必须drop目标列 X = df.drop(columns=['PRZYPADKOWE_CZY_CELOWE']) y = df['PRZYPADKOWE_CZY_CELOWE'].astype(int) # 强制转整数,避免隐式类型转换 # 2. 拆分数据集,分层抽样保证标签分布一致 X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) # 3. 配置预处理逻辑 preprocessor = ColumnTransformer( transformers=[ ('num', StandardScaler(), make_column_selector(dtype_include=np.number)), ('cat', OneHotEncoder(handle_unknown='ignore'), make_column_selector(dtype_include=object)) ] ) # 4. 构建端到端Pipeline model = make_pipeline( preprocessor, LogisticRegression(max_iter=1000, random_state=42) ) # 5. 训练、计算预测得分 model.fit(X_train, y_train) y_score = model.decision_function(X_test) # 6. 计算ROC指标并绘图 fpr, tpr, thresholds = roc_curve(y_true=y_test, y_score=y_score, pos_label=1) RocCurveDisplay(fpr=fpr, tpr=tpr).plot()
其他易触发报错的场景
- 数值列存在隐式字符串缺失值(比如填了
'N/A'、'缺失'这类字符串),导致列类型被识别为object,被OneHotEncoder编码后特征维度错位,需要先对数值列做类型清洗,把非法字符串值替换为np.nan后再做缺失值填充。 - 调用
roc_curve时没有按参数顺序传值,把y_score和y_true位置写反,导致函数拿字符串类型的标签数组和阈值做大小比较触发报错。
内容的提问来源于stack exchange,提问作者PawelKinczyk
相关产品推荐
相关产品推荐

