
scikit-learn 交叉验证与数组 API混合命名空间下y处理策略的修复解析【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn导读本文聚焦 scikit-learn 数组 APIArray API支持中的一个重要修复GridSearchCV、RandomizedSearchCV、cross_validate与cross_val_score等交叉验证工具在输入为数组 API 兼容的X如 CuPy、PyTorch 张量时不再强制把目标值y转换到X的数组命名空间而是把y的命名空间与设备处理完全委托给被包装的估计器。通过阅读本文你将理解这一改动背后的设计动机、它在源码中的落点索引切分与命名空间探测的分离以及如何组合「数组 APIX NumPyy含字符串y」进行交叉验证。该条目来自 scikit-learn 仓库的更新日志 doc/whats_new/upcoming_changes/array-api/33633.fix.rst由 Tim Head 贡献属于数组 API 支持方向的缺陷修复fix。一、背景scikit-learn 的数组 API 调度机制在展开这个修复之前需要先理解 scikit-learn 如何支持数组 API。仓库中sklearn/utils/_array_api.py的get_namespace函数定义于 L379-L471负责探测输入数组所属的命名空间它内部调用array_api_compat.get_namespace(*arrays)来获得多个数组共享的数组 API 兼容命名空间只有sklearn.set_config(array_api_dispatchTrue)或sklearn.config_context(array_api_dispatchTrue)上下文显式开启时数组 API 调度才会生效否则始终回退到 NumPy 命名空间返回值为(namespace, is_array_api_compliant)其中is_array_api_compliant表示数组是否实现了__array_namespace__协议即符合 NEP-47 数组 API 标准。关键点在于get_namespace的注释明确写道——If any of thearraysare not arrays, the namespace defaults to the NumPy namespace只要传入的某个对象不是数组命名空间就回退为 NumPy。这也正是本次修复能成立的底层原因之一当y是普通 NumPy 数组甚至字符串数组时不应该强迫它参与数组 API 命名空间探测。二、修复内容y不再被强制拉入X的命名空间2.1 修复前的问题在旧实现中交叉验证入口会对X和y一视同仁地做命名空间统一处理。当用户传入一个数组 API 的X例如位于 GPU 上的 CuPy 数组和一个普通的 NumPyy时交叉验证工具会尝试把y也转换到X的数组命名空间。这带来两个实际问题不必要的拷贝/转换y被强行搬到X所在的设备与命名空间增加开销无法处理字符串y数组 API 命名空间如 PyTorch、CuPy往往不支持字符串 dtype。如果y是np.array([a, b])这样的 NumPy 字符串数组它根本无法被数组 API 命名空间持有强制转换会直接报错。2.2 修复后的行为本次修复PR #33633后GridSearchCV、RandomizedSearchCV、cross_validate、cross_val_score这四个交叉验证工具不再在交叉验证过程中强制把y转换到X的数组命名空间y的命名空间namespace与设备device处理委托给被包装的估计器the wrapped estimator由估计器自身的fit在真正消费y时自行决定如何处理因此数组 API 的X可以与 NumPyy自由组合包括 NumPy 字符串y——这类y只需在估计器内部保持 NumPy 语义即可无需进入数组 API 命名空间。2.3 适用对象速览工具所属模块行为变化GridSearchCVsklearn.model_selection._searchfit中y不再被转换命名空间RandomizedSearchCVsklearn.model_selection._search同上与GridSearchCV共享BaseSearchCV.fitcross_validatesklearn.model_selection._validationX, y indexable(X, y)后原样传给估计器cross_val_scoresklearn.model_selection._validation内部调用cross_validate行为一致三、源码落点命名空间探测与索引切分的解耦本次修复的实质是让交叉验证入口只做样本索引切分不做命名空间归一化。这一点可以从两条关键调用链得到印证。3.1 入口处indexable仅校验、不转换在sklearn/model_selection/_validation.py的cross_validate中L318X, y indexable(X, y)在sklearn/model_selection/_search.py的BaseSearchCV.fit中L1000X, y indexable(X, y)indexable的作用仅是确认输入支持__len__、整数索引与切片从而保证后续能按折切分它不会把y转换到任何数组 API 命名空间。也就是说y在进入交叉验证循环前保持用户传入时的原始类型NumPy 数组、字符串数组均可。3.2 切分处_safe_split只索引、不改写在_fit_and_score内部每个折的训练/测试子集通过_safe_split生成L829-L830X_train, y_train _safe_split(estimator, X, y, train) X_test, y_test _safe_split(estimator, X, y, test, train)_safe_split定义在sklearn/utils/metaestimators.py其对y的处理非常简单L172-L175if y is not None: y_subset _safe_indexing(y, indices) else: y_subset None这里仅通过_safe_indexing按折索引y自始至终没有调用get_namespace或任何数组 API 转换函数。因此如果X是 CuPy/PyTorch 数组X_subset仍是数组 API 对象切分后的X_train/X_test自然留在原命名空间与设备上y_subset则保持 NumPy 语义是否迁移命名空间、迁移到哪个设备完全交给估计器fit(X_train, y_train)内部自行决定。这正是 changelog 所述 Namespace and device handling ofyis delegated to the wrapped estimator 的源码实现。四、测试验证混合输入的官方保障仓库为这次修复提供了专门的参数化测试位于sklearn/model_selection/tests/test_validation.py的test_cross_validate_array_api_mixed_inputspytest.mark.parametrize(y_is_string, [False, True]) pytest.mark.parametrize( array_namespace, device_name, dtype_name, yield_namespace_device_dtype_combinations(), ) def test_cross_validate_array_api_mixed_inputs( array_namespace, device_name, dtype_name, y_is_string ): Check cross_validate works with array API X and NumPy y. xp, device _array_api_for_tests(array_namespace, device_name) X_np np.arange(100).reshape((10, 10)).astype(dtype_name) X_xp xp.asarray(X_np, devicedevice) y_np np.array([0] * 5 [1] * 5) if y_is_string: y_np np.array([a, b])[y_np] with config_context(array_api_dispatchTrue): cross_validate( LogisticRegression(), X_xp, y_np, cv2, error_scoreraise, )该测试揭示了修复后的三个关键保证组合合法性X为数组 API 数组由_array_api_for_tests在多种命名空间/设备组合下创建y为普通 NumPy 数组cross_validate必须正常运行字符串y场景y_is_stringTrue时y是np.array([a, b])[y_np]产生的字符串数组——这是数组 API 命名空间无法持有的类型修复后依然可以顺利走完整个交叉验证流程显式开启调度测试在config_context(array_api_dispatchTrue)下运行确保数组 API 调度确实被激活从而真实覆盖「X走数组 API、y走 NumPy」的混合路径。此外测试注释也点明了该用例的必要性cross_validateis a function so it is not covered by the commoncheck_array_api_*estimator checks——即函数形式的工具不在通用估计器数组 API 检查check_array_api_*覆盖范围内必须单独编写测试守护。与之配套的还有sklearn/model_selection/tests/test_search.py中对SearchCVGridSearchCV/RandomizedSearchCV的数组 API 测试以及sklearn/model_selection/tests/test_split.py中与切分器相关的命名空间/设备参数化用例共同保障了交叉验证全链路在混合输入下的正确性。五、实际使用示例与注意事项5.1 可直接运行的示例import numpy as np from sklearn.datasets import load_iris from sklearn.linear_model import LogisticRegression from sklearn.model_selection import cross_validate, GridSearchCV X, y load_iris(return_X_yTrue) # 将 X 转为数组 API 数组此处以 CuPy 为例需已安装 cupy import cupy as cp X_xp cp.asarray(X) # 字符串 y数组 API 命名空间无法持有但修复后可以混用 y_str np.array([setosa, versicolor, virginica])[y] # 1) cross_validate 直接混用 import sklearn with sklearn.config_context(array_api_dispatchTrue): results cross_validate( LogisticRegression(max_iter1000), X_xp, y_str, cv3, error_scoreraise ) print(results[test_score]) # 2) GridSearchCV 同样支持 with sklearn.config_context(array_api_dispatchTrue): search GridSearchCV( LogisticRegression(max_iter1000), {C: [0.1, 1.0, 10.0]}, cv3, error_scoreraise, ) search.fit(X_xp, y_str) print(search.best_params_)5.2 注意事项必须显式开启数组 API 调度如sklearn/utils/_array_api.py文档所示数组 API 支持默认关闭。只有执行sklearn.set_config(array_api_dispatchTrue)或使用sklearn.config_context(array_api_dispatchTrue)上下文X才会走数组 API 路径否则一切输入按 NumPy 处理。y的类型决定权在估计器修复后交叉验证工具不干预y这意味着最终y是否被迁移到X的设备/命名空间取决于被包装估计器fit的实现。对于官方估计器y通常会被安全地转换为训练所需类型但第三方自定义估计器需要自行保证对混合命名空间输入的兼容。groups与 fit 参数走独立路径groups传给切分器splitter用于决定划分方式而不是作为y参与命名空间处理长度与样本数相同的 fit 参数如sample_weight则仍会按折切分参见sklearn/model_selection/_search.py对fit参数的说明。它们的命名空间处理不受本次修复影响。六、小结本次修复33633.fix.rst是 scikit-learn 数组 API 支持演进中的一个务实修正交叉验证基础设施回归其「只管切分、不管转换」的职责边界将y的命名空间与设备决策下沉到估计器层面。其直接收益是解锁了两类此前无法工作的组合——数组 APIX与 NumPyy、数组 APIX与 NumPy 字符串y同时消除了对y的无谓转换开销。从源码看这一语义由indexable_safe_split的「只索引不转换」链路保证并由test_cross_validate_array_api_mixed_inputs等参数化测试持续守护。【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考