news 2026/9/18 14:32:27

scikit-learn Array API 增强:分类指标函数支持跨命名空间、跨设备的混合数组输入

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
scikit-learn Array API 增强:分类指标函数支持跨命名空间、跨设备的混合数组输入

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_curveroc_curvezero_one_lossjaccard_scorebalanced_accuracy_scorecohen_kappa_score这六个分类/排序指标函数如何在 Array API 分发(dispatch)下接受"来自混合命名空间与混合设备"的数组输入。文中结合 sklearn/metrics 与 sklearn/utils/_array_api.py 的源码,说明混合输入被自动对齐的机制、测试验证方式,以及启用该特性时的前置条件与注意事项。读完后你可以掌握:如何在 GPU 分数数组与 CPU 标签数组并存的场景下调用这些指标函数,以及"everything follows y_pred"这一转换规则是如何在底层落地的。

1. 增强内容概览:六个指标函数获得混合输入支持

该增强条目(贡献者 Lucy Liu)的原文要点是:

  • 为以下函数新增"来自混合命名空间和混合设备"的数组输入支持:
    • sklearn.metrics.det_curve
    • sklearn.metrics.roc_curve
    • sklearn.metrics.zero_one_loss
    • sklearn.metrics.jaccard_score
    • sklearn.metrics.balanced_accuracy_score
    • sklearn.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_validateGridSearchCV等 meta-estimator 内部调用的评分函数就会同时收到 GPU 上的y_pred和 CPU 上的y_true。只有评分函数支持混合输入,这类跨设备的评估流程才能跑通。

2. 使用前提:启用 array_api_dispatch

Array API 支持在 scikit-learn 中仍被标记为实验性,需要显式开启,且要求安装较新版本的依赖。要点如下(详见 doc/modules/array_api.rst):

  1. 全局或局部开启分发:
from sklearn import config_context, set_config # 方式一:全局开启(官方推荐,避免意外混合数组命名空间) set_config(array_api_dispatch=True) # 方式二:临时上下文管理器,退出 with 块后自动恢复 with config_context(array_api_dispatch=True): ...
  1. 在导入scipyscikit-learn之前设置环境变量SCIPY_ARRAY_API=1,以启用 SciPy 自身的 Array API 支持。源码中这一前置条件由_check_array_api_dispatch强制校验:SciPy 版本低于 1.14 会抛ImportError,环境变量未设置会抛RuntimeError,见 sklearn/utils/_array_api.py#L168。

  2. 分发关闭时的行为差异: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是基准,ysample_weight等其余数组输入全部被转换到X的库与设备;
  • 评分函数(scoring functions)y_pred(曲线类函数中即y_score)是基准,y_truesample_weight等被转换到y_pred的库与设备。

本次增强涉及的六个函数均属评分函数,因此遵循第二条规则。输出的类型约定为:

  • 返回标量的函数(如zero_one_lossjaccard_scorebalanced_accuracy_scorecohen_kappa_score)返回 Python 标量(通常是float),而非数组标量;
  • 返回数组的函数(如roc_curvedet_curve)返回与y_pred同库、同设备的数组。

此外还有一个特殊情形:混合输入支持也覆盖了"y_true是 NumPy 字符串数组、其余输入是任意容器类型的数值数组"的情况。由于数组 API 规范只覆盖数值数组,scikit-learn 会把y先转成数值表示(如 one-hot / 序数编码),再移动到其余输入的命名空间与设备上。

3.1 底层实现:get_namespace_and_devicemove_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_curveroc_curve

两者都位于 sklearn/metrics/_ranking.py:

  • det_curve(函数定义):检测错误权衡曲线,仅支持二分类任务。参数为y_truey_scorepos_label=Nonesample_weight=Nonedrop_intermediate=False。返回fprfnrthresholds三个数组。实现上先调用confusion_matrix_at_thresholds统计各阈值下的真/假阳性,再用xp.concat在头部追加"阈值取无穷大、恒预测负类"的端点(1.7 版起的行为),并支持drop_intermediate剔除tp不变的中间阈值点。混合输入下,y_score决定命名空间与设备,y_true/sample_weightmove_to对齐过去,三个输出数组与y_score同库同设备。
  • roc_curve(函数定义):接收同样的参数集,返回fprtprthresholds。它与det_curve共享confusion_matrix_at_thresholds的统计路径,混合输入的处理方式一致。

4.2 标量类:zero_one_lossjaccard_scorebalanced_accuracy_scorecohen_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源码交集/并集;支持多标签与多类,常用参数labelspos_labelaveragesample_weight标量或每类分数数组
balanced_accuracy_score源码balanced_accuracy_score(y_true, y_pred, *, sample_weight=None, adjusted=False)各类召回率的均值,Python float
cohen_kappa_score源码常用参数labelsweightssample_weightadjust_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. 测试如何验证混合输入行为

仓库内有两层测试保障:

  1. 混合输入参数组合的挑选: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 cudanumpy → torch cudatorch mps → numpy等。
  2. 结果一致性检查: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" -v

array-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=TrueSCIPY_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 14:31:07

为什么 coding agent 主流选择 Node.js 而非 Rust 或 Python

1. 为什么市面上的 coding agent 大多数都基于 Node.js&#xff1f;——一个从业十年的全栈工程师的硬核拆解你打开 GitHub Trending&#xff0c;刷一遍最近三个月爆火的 coding agent 项目&#xff1a;Cursor、Tabby、Continue、Bloop、CodeWhisperer 的开源替代品、甚至不少大…

作者头像 李华
网站建设 2026/9/18 14:31:07

PI-Desktop架构全解:Electron、Rust Host Core与pi Agent Sidecar的分工

PI-Desktop架构全解&#xff1a;Electron、Rust Host Core与pi Agent Sidecar的分工 【免费下载链接】PI-Desktop Local-first AI coding agent desktop: Electron Rust host core pi Agent Harness user-installable plugins 项目地址: https://gitcode.com/GitHub_Trend…

作者头像 李华
网站建设 2026/9/18 14:29:58

Python+OpenGL绘制3D模型(五)绘制三角型

系列文章 基础 PythonOpenGL绘制3D模型&#xff08;一&#xff09;Python 和 PyQt环境搭建 PythonOpenGL绘制3D模型&#xff08;二&#xff09;程序框架PyQt5 PythonOpenGL绘制3D模型&#xff08;三&#xff09;程序框架PyQt6 PythonOpenGL绘制3D模型&#xff08;四&#xff0…

作者头像 李华
网站建设 2026/9/18 14:29:38

给 docmd 的 Markdown 文档站加 MCP,TaoToken 只提供 Key

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华