scikit-learn Array API 增强:分类指标函数支持跨命名空间、跨设备的混合数组输入
【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn
本文围绕 scikit-learn 增强条目 34442.enhancement.rst 展开,讲解det_curve、roc_curve、zero_one_loss、jaccard_score、balanced_accuracy_score、cohen_kappa_score这六个分类/排序指标函数如何在 Array API 分发(dispatch)下接受"来自混合命名空间与混合设备"的数组输入。文中结合 sklearn/metrics 与 sklearn/utils/_array_api.py 的源码,说明混合输入被自动对齐的机制、测试验证方式,以及启用该特性时的前置条件与注意事项。读完后你可以掌握:如何在 GPU 分数数组与 CPU 标签数组并存的场景下调用这些指标函数,以及"everything follows y_pred"这一转换规则是如何在底层落地的。
1. 增强内容概览:六个指标函数获得混合输入支持
该增强条目(贡献者 Lucy Liu)的原文要点是:
- 为以下函数新增"来自混合命名空间和混合设备"的数组输入支持:
sklearn.metrics.det_curvesklearn.metrics.roc_curvesklearn.metrics.zero_one_losssklearn.metrics.jaccard_scoresklearn.metrics.balanced_accuracy_scoresklearn.metrics.cohen_kappa_score
这里的两个关键词需要拆开理解:
- 命名空间(namespace):实现 Array API 规范的数组库。scikit-learn 定期做合规性测试的库包括 PyTorch(CPU/CUDA/MPS/XPU)、CuPy(CUDA)、dpnp(CPU/Intel GPU),见 doc/modules/array_api.rst。
- 设备(device):同一命名空间内数据驻留的硬件位置,例如 torch 的
"cpu"、"cuda"、"mps",dpnp 的"cpu"、"gpu"。
在此之前,启用 Array API 分发时,一个指标函数的所有数组输入通常需要处于同一命名空间且同一设备。本次增强后,上述六个函数允许输入来自不同的库、不同的设备(例如y_score是 torch CUDA 张量而y_true是 NumPy 数组),scikit-learn 会在内部自动把不一致的输入转换到统一的命名空间与设备上。
这项能力的实际价值在于:在流水线中X可能被FunctionTransformer之类的步骤移动到 GPU 以提升性能,而y因 scikit-learn 的 Pipeline 不允许对y做变换(避免数据泄漏)仍留在 CPU。此时cross_validate、GridSearchCV等 meta-estimator 内部调用的评分函数就会同时收到 GPU 上的y_pred和 CPU 上的y_true。只有评分函数支持混合输入,这类跨设备的评估流程才能跑通。
2. 使用前提:启用 array_api_dispatch
Array API 支持在 scikit-learn 中仍被标记为实验性,需要显式开启,且要求安装较新版本的依赖。要点如下(详见 doc/modules/array_api.rst):
- 全局或局部开启分发:
from sklearn import config_context, set_config # 方式一:全局开启(官方推荐,避免意外混合数组命名空间) set_config(array_api_dispatch=True) # 方式二:临时上下文管理器,退出 with 块后自动恢复 with config_context(array_api_dispatch=True): ...在导入
scipy与scikit-learn之前设置环境变量SCIPY_ARRAY_API=1,以启用 SciPy 自身的 Array API 支持。源码中这一前置条件由_check_array_api_dispatch强制校验:SciPy 版本低于 1.14 会抛ImportError,环境变量未设置会抛RuntimeError,见 sklearn/utils/_array_api.py#L168。分发关闭时的行为差异:
array_api_dispatch=False时,所有 array-like 输入会用numpy.asarray转成 NumPy 数组,输出也一定是 NumPy;而 GPU 上的 torch 张量通常无法被转成 NumPy,即直接报错。这也是官方建议在 Array API 输入场景下始终开启分发的原因。array_api_dispatch=True时,输出数组的库与设备取决于输入(见下文第 3 节规则)。
3. 核心规则:"everything follows y_pred"
doc/modules/array_api.rst 明确定义了两套对齐规则:
- 对估计器(estimators):
X是基准,y、sample_weight等其余数组输入全部被转换到X的库与设备; - 对评分函数(scoring functions):
y_pred(曲线类函数中即y_score)是基准,y_true、sample_weight等被转换到y_pred的库与设备。
本次增强涉及的六个函数均属评分函数,因此遵循第二条规则。输出的类型约定为:
- 返回标量的函数(如
zero_one_loss、jaccard_score、balanced_accuracy_score、cohen_kappa_score)返回 Python 标量(通常是float),而非数组标量; - 返回数组的函数(如
roc_curve、det_curve)返回与y_pred同库、同设备的数组。
此外还有一个特殊情形:混合输入支持也覆盖了"y_true是 NumPy 字符串数组、其余输入是任意容器类型的数值数组"的情况。由于数组 API 规范只覆盖数值数组,scikit-learn 会把y先转成数值表示(如 one-hot / 序数编码),再移动到其余输入的命名空间与设备上。
3.1 底层实现:get_namespace_and_device与move_to
转换机制集中在 sklearn/utils/_array_api.py:
get_namespace(*arrays)(源码):内省数组参数,返回其共同的 Array API 命名空间对象;对普通 NumPy 数组返回array_api_compat.numpy包装。get_namespace_and_device(*arrays)(源码):在上面基础上再提取数组驻留的硬件设备,供调用方确定"以谁为基准"。move_to(*arrays, xp, device)(源码):把数组移动到基准命名空间与设备。跨命名空间/跨设备转移时优先尝试 DLPack 协议(xp.from_dlpack(array, device=device),零拷贝且库无关),在目标库不支持 DLPack 1.0(AttributeError/TypeError/NotImplementedError等)时回退到"经 NumPy 中转"的两步转换(A → numpy → B)。若目标设备不支持float64(如 MPS、部分 XPU 设备),float64数组会被降精度到float32。
以det_curve为例,其入口第一行即为xp, _, device = get_namespace_and_device(y_score)(sklearn/metrics/_ranking.py#L407),随后所有后续xp.*运算都在该命名空间与设备上执行;y_true的命名空间仅在需要时单独内省(sklearn/metrics/_ranking.py#L439),并不参与基准判定——这正是"y_pred 说了算"的直接体现。
4. 六个函数的源码定位与行为说明
4.1 曲线类:det_curve与roc_curve
两者都位于 sklearn/metrics/_ranking.py:
det_curve(函数定义):检测错误权衡曲线,仅支持二分类任务。参数为y_true、y_score、pos_label=None、sample_weight=None、drop_intermediate=False。返回fpr、fnr、thresholds三个数组。实现上先调用confusion_matrix_at_thresholds统计各阈值下的真/假阳性,再用xp.concat在头部追加"阈值取无穷大、恒预测负类"的端点(1.7 版起的行为),并支持drop_intermediate剔除tp不变的中间阈值点。混合输入下,y_score决定命名空间与设备,y_true/sample_weight被move_to对齐过去,三个输出数组与y_score同库同设备。roc_curve(函数定义):接收同样的参数集,返回fpr、tpr、thresholds。它与det_curve共享confusion_matrix_at_thresholds的统计路径,混合输入的处理方式一致。
4.2 标量类:zero_one_loss、jaccard_score、balanced_accuracy_score、cohen_kappa_score
四者都位于 sklearn/metrics/_classification.py:
| 函数 | 位置 | 参数要点 | 返回值 |
|---|---|---|---|
zero_one_loss | 源码 | zero_one_loss(y_true, y_pred, *, normalize=True, sample_weight=None) | 误判比例(normalize=True时为 0-1 之间的 Python float) |
jaccard_score | 源码 | 交集/并集;支持多标签与多类,常用参数labels、pos_label、average、sample_weight | 标量或每类分数数组 |
balanced_accuracy_score | 源码 | balanced_accuracy_score(y_true, y_pred, *, sample_weight=None, adjusted=False) | 各类召回率的均值,Python float |
cohen_kappa_score | 源码 | 常用参数labels、weights、sample_weight、adjust_to_pooled | 超出偶然一致程度的 Kappa 系数,Python float |
这些函数内部的标签校验(_check_targets,源码)与混淆矩阵构建(confusion_matrix,源码)同样做了命名空间分发,因此y_true(含 NumPy 字符串标签这一特例)与y_pred设备不一致时,会被自动搬运到y_pred的命名空间与设备上再计算。按第 3 节的输出约定,标量结果直接是 Pythonfloat,不产生设备上的数组标量。
5. 实战示例:NumPy 标签 + PyTorch GPU 分数
下面的示例演示六个函数在混合输入下的调用形态。前提是机器具备 CUDA、已安装 PyTorch 与新版 SciPy(并先设置SCIPY_ARRAY_API=1);没有 GPU 时可用array-api-strict的模拟设备做等价验证(见第 6 节):
# 必须在 import scipy / sklearn 之前设置 import os os.environ["SCIPY_ARRAY_API"] = "1" import numpy as np import torch from sklearn import config_context from sklearn.metrics import ( balanced_accuracy_score, c Cohen_kappa_score, det_curve, jaccard_score, roc_curve, zero_one_loss, ) # y_true / y 在 CPU 的 NumPy 上(模拟流水线中 y 未被移动到 GPU 的情形) y_true_np = np.array([0, 0, 1, 1, 0, 1]) y_pred_np = np.array([0, 0, 1, 1, 1, 0]) # y_score / 概率分数在 CUDA 的 PyTorch 上 y_score_t = torch.tensor([0.1, 0.4, 0.35, 0.8, 0.6, 0.9], device="cuda") with config_context(array_api_dispatch=True): # 曲线类:输出数组跟随 y_score 的命名空间与设备 fpr, tpr, thr = roc_curve(y_true_np, y_score_t) fpr_det, fnr_det, thr_det = det_curve(y_true_np, y_score_t, drop_intermediate=True) print(type(fpr), fpr.device) # <class 'torch.Tensor'> cuda print(type(fpr_det), fpr_det.device) # <class 'torch.Tensor'> cuda # 标量类:y_true/y 为 NumPy,y_pred 为 CUDA 张量,返回 Python float print(zero_one_loss(y_true_np, torch.asarray(y_pred_np, device="cuda"))) print(jaccard_score(y_true_np, torch.asarray(y_pred_np, device="cuda"))) print(balanced_accuracy_score(y_true_np, torch.asarray(y_pred_np, device="cuda"))) print(c Cohen_kappa_score(y_true_np, torch.asarray(y_pred_np, device="cuda")))若同一输入在array_api_dispatch=False下运行,y_score_t这类 GPU 张量在numpy.asarray转换阶段就会失败——这正是混合输入支持的用武之地。
6. 测试如何验证混合输入行为
仓库内有两层测试保障:
- 混合输入参数组合的挑选:sklearn/utils/_array_api.py#L120 的
yield_mixed_namespace_input_permutations定义了测试用的(输入命名空间/设备 → 参考命名空间/设备)组合,覆盖非 NumPy→NumPy(GPU→CPU)、NumPy→非 NumPy(CPU→GPU)、非 NumPy→非 NumPy(GPU→GPU)以及array-api-strict→非 NumPy(本地无硬件即可跑)四类转换方向,例如cupy → torch cuda、numpy → torch cuda、torch mps → numpy等。 - 结果一致性检查:sklearn/metrics/tests/test_common.py#L2700 附近的通用检查会构造混合命名空间输入,断言"输出与纯 NumPy 参考结果一致、且
y_true/sample_weight跟随y_pred",错误信息形如 "Output incorrect for mixed namespace and device array input to ..."(相关断言);字符串y与数值输入混合的场景另有专门检查(源码)。估计器侧的等价检查位于 sklearn/utils/estimator_checks.py#L1434。
开发者本地不需要 GPU 即可回归这些路径:安装array-api-strict后执行
pip install array-api-strict pytest -k "array_api" -varray-api-strict提供带模拟设备的严格 Array API 实现,能快速暴露多设备处理问题;对真实 CUDA/MPS/Intel GPU 硬件的覆盖则由 CI 在 pull request 上执行,无法执行的检查会自动跳过,因此建议加-v观察跳过项。
7. 注意事项与适用范围
- 实验性状态:Array API 分发属于实验特性,官方不做向后兼容承诺;依赖库版本过旧可能不工作(详见 doc/modules/array_api.rst 的"Enabling array API support"一节)。
- 基准是 y_pred:这六个函数中,
y_score/y_pred决定输出的库与设备;不要把期望的"目标设备"寄托在y_true上。 - float64 的设备限制:在 PyTorch MPS 与部分 Intel GPU 设备上不支持
float64,scikit-learn 会自动回退到float32,可能与 CPU 路径在数值上不一致(见 doc/modules/array_api.rst 中"Note on device support for float64")。 - 与既有支持面的关系:这六个函数本就已列入 Array API 支持指标清单(doc/modules/array_api.rst 的 Metrics 一节),本次增强是把它们从"要求所有输入同库同设备"扩展为"接受混合库/混合设备并自动对齐";完整支持矩阵(估计器、meta-estimators、工具函数)以该文档为准。
- 关闭分发时:一切输入经
numpy.asarray落到 NumPy,GPU 数组无法参与计算,这是使用混合输入示例的硬性前提——必须同时满足array_api_dispatch=True与SCIPY_ARRAY_API=1。
综上,这条增强以最小的用户侧改动(无需手动搬运y_true),补齐了跨设备模型选择与评估链路上最后一块拼图:评分函数。其可验证的依据是 sklearn/metrics/_ranking.py 与 sklearn/metrics/_classification.py 中的命名空间分发实现,以及 sklearn/utils/_array_api.py 提供的get_namespace_and_device/move_to基础设施与 sklearn/metrics/tests/test_common.py 中的混合输入回归检查。
【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考