scikit-learn 全面拥抱 Array API 标准:GPU 加速与混合设备计算时代开启
更新背景
Python 数据 API 标准联盟(Consortium for Python Data API Standards)推出的 Array API 标准,旨在为各类数组库定义一致接口,使“数组消费型”库(如 scikit-learn)能够编写与底层数组实现无关的代码。这一变革对 scikit-learn 而言具有里程碑意义:它解决了困扰该工具长达 11 年的 GPU 支持难题。过去,由于软件依赖复杂及平台特异性问题,scikit-learn 曾明确表示短期内不会添加 GPU 支持。如今,依托 Array API 标准,这些障碍已被彻底消除,用户可无缝利用 PyTorch、CuPy 等库的硬件加速能力。
核心突破与功能特性
1. 内建 Vendor 库与标准化兼容
scikit-learn 现已内建(vendor)成熟的 array-api-compat 和 array-api-extra 库:
- array-api-compat:作为 PyTorch、CuPy、JAX 等库的包装器,填补标准与具体实现间的差距,确保向后兼容。
- array-api-extra:提供标准之外但对数组消费库至关重要的扩展函数。 此举避免了代码库中复杂的条件依赖处理,遵循了 SciPy 的最佳实践。
2. 广泛的后端支持与设备覆盖
当前已支持以下数组库及设备:
- CuPy:完整的 ndarray 支持。
- PyTorch:覆盖 CPU、CUDA、MPS(Apple Silicon)及 XPU(Intel)所有设备。
- NumPy:作为基础支持。
- JAX:支持正在推进中。
此外,scikit-learn 还通过
array-api-strict进行严格合规性测试,确保符合标准的库无需额外修改即可被接受。
3. 混合数组命名空间与设备处理
这是 scikit-learn 独有的架构设计,允许特征(X)与标签(y)使用不同数组库或设备:
- 场景示例:字符串类别标签通常仅由 NumPy 支持,而计算需利用 GPU。该架构允许在 CPU 上对字符串标签进行编码(如
TargetEncoder),同时通过FunctionTransformer将特征数组转换为 CUDA 张量,送入RidgeClassifier进行 GPU 加速训练。 - 流水线(Pipeline)增强:解决了传统流水线无法修改目标数组(y)的限制,使得混合输入成为可能,极大提升了端到端工作流的灵活性。
4. 关键模型与指标支持
大量高影响力指标(Metrics)和转换器(Transformers)已适配,包括 LabelBinarizer。复杂的估计器(Estimators)也已完成多项核心模型的重构,包括:
LogisticRegressionGaussianNB,GaussianMixtureRidge及其变体(RidgeCV,RidgeClassifier等)Nystroem,PCAGaussianProcessRegressor(开发进行中)
实际应用价值
对于开发者而言,此次更新意味着无需编写复杂的后端适配代码,即可构建跨硬件的机器学习管道。对于企业用户,这意味着能够直接利用企业级 GPU 集群进行大规模模型训练,同时保留 scikit-learn 熟悉的 API 风格。混合设备支持尤其适合处理包含文本标签(需 CPU 处理)和数值特征(需 GPU 加速)的复杂数据集,显著降低了异构计算的工作门槛。
“通过 Array API 标准,我们不仅解决了 GPU 支持的长期悬而未决的问题,更重新定义了 scikit-learn 在异构计算环境中的角色。” —— Lucy Liu, scikit-learn 团队
技术展望
随着 JAX 支持的临近以及更多估计器的完成,scikit-learn 正逐步构建一个真正无边界的数据科学生态。未来,随着更多库遵循 Array API 标准,该工具将成为连接传统机器学习与现代高性能计算的关键枢纽。
注:具体支持的指标与估计器列表请参阅官方文档中的 Array API support 页面。