手动交叉验证与cross_val_score中Sklearn precision_score行为不一致
precision_score设置zero_division=np.nan时手动交叉验证正常,但cross_val_score报错
我在precision_score中设置zero_division参数为np.nan,手动执行交叉验证时运行正常,但使用cross_val_score时触发参数验证错误。
可复现场景
(附可复现数据文件)
# 加载数据 import pickle import numpy as np from sklearn.metrics import precision_score from sklearn.model_selection import cross_val_score with open("sklearn_data.pkl", "rb") as f: objects = pickle.load(f) # objects包含:estimator, X, y, scoring, cv, n_jobs estimator = objects["estimator"] X = objects["X"] y = objects["y"] scoring = objects["scoring"] # 配置为make_scorer(precision_score, pos_label="Case_0", zero_division=np.nan) cv = objects["cv"] n_jobs = objects["n_jobs"] # 验证所有训练/验证集都包含两类样本 pos_label = "Case_0" control_label = "Control" for index_training, index_validation in cv: assert y.iloc[index_training].nunique() == 2 assert y.iloc[index_validation].nunique() == 2 assert pos_label in y.values assert control_label in y.values # 手动交叉验证:运行正常,输出均值0.501156937317928 scores = list() for index_training, index_validation in cv: estimator.fit(X.iloc[index_training], y.iloc[index_training]) y_hat = estimator.predict(X.iloc[index_validation]) score = precision_score(y_true=y.iloc[index_validation], y_pred=y_hat, pos_label=pos_label) scores.append(score) print(np.mean(scores)) # 使用cross_val_score:报错 cross_val_score(estimator=estimator, X=X, y=y, cv=cv, scoring=scoring, n_jobs=n_jobs)
报错信息
/Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/python3.9/site-packages/sklearn/model_selection/_validation.py:839: UserWarning: Scoring failed. The score on this train-test partition for these parameters will be set to nan. Details: Traceback (most recent call last): File "/Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/python3.9/site-packages/sklearn/metrics/_scorer.py", line 136, in __call__ score = scorer._score( File "/Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/python3.9/site-packages/sklearn/metrics/_scorer.py", line 355, in _score return self._sign * self._score_func(y_true, y_pred, **scoring_kwargs) File "/Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/python3.9/site-packages/sklearn/utils/_param_validation.py", line 201, in wrapper validate_parameter_constraints( File "/Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/python3.9/site-packages/sklearn/utils/_param_validation.py", line 95, in validate_parameter_constraints raise InvalidParameterError( sklearn.utils._param_validation.InvalidParameterError: The 'zero_division' parameter of precision_score must be a float among {0.0, 1.0, nan} or a str among {'warn'}. Got nan instead.
环境版本
系统信息: python: 3.9.16 | packaged by conda-forge | (main, Feb 1 2023, 21:42:20) [Clang 14.0.6 ] 可执行文件路径: /Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/bin/python 机器环境: macOS-13.4.1-x86_64-i386-64bit Python依赖包版本: sklearn: 1.3.1 pip: 22.0.3 setuptools: 60.7.1 numpy: 1.24.4 scipy: 1.8.0 Cython: 0.29.27 pandas: 1.4.0 matplotlib: 3.7.1 joblib: 1.3.2 threadpoolctl: 3.1.0 基于OpenMP构建: True threadpoolctl信息: user_api: blas internal_api: openblas prefix: libopenblas filepath: /Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/libopenblasp-r0.3.18.dylib version: 0.3.18 threading_layer: openmp architecture: Haswell num_threads: 16 user_api: openmp internal_api: openmp prefix: libomp filepath: /Users/jespinoz/anaconda3/envs/soothsayer_py3.9_env2/lib/libomp.dylib version: None num_threads: 16
解决方案
问题出在np.nan与scikit-learn参数验证的兼容性上:当cross_val_score启用多进程(n_jobs>1)时,scorer对象会被序列化,np.nan在序列化后可能被识别为非预期的类型,触发参数验证失败。
两种解决办法:
改用Python原生
float('nan')
创建scorer时替换np.nan为float('nan'):from sklearn.metrics import make_scorer, precision_score scoring = make_scorer(precision_score, pos_label="Case_0", zero_division=float('nan'))自定义scorer函数封装
把precision_score的调用封装在自定义函数中,避免直接传递np.nan作为参数:def custom_precision(y_true, y_pred): return precision_score(y_true, y_pred, pos_label="Case_0", zero_division=np.nan) scoring = make_scorer(custom_precision)
原因说明
scikit-learn的参数验证逻辑中,zero_division允许的值包含nan,但这里的nan指Python原生的float('nan')。np.nan虽然值上等于原生nan,但属于numpy的float类型,在多进程序列化过程中可能被转换,导致参数验证时无法匹配预期的类型约束。手动交叉验证没有序列化步骤,直接调用precision_score时numpy的nan可以被兼容,因此不会报错。
内容的提问来源于stack exchange,提问作者O.rka
相关产品推荐
相关产品推荐

