11.1. Поддержка Array API (экспериментальная)
Спецификация Array API определяет стандартный API для всех библиотек работы с массивами, имеющих API наподобие NumPy. Поддержка Array API в Scikit-learn требует установки array-api-compat, а также переменная среды SCIPY_ARRAY_API должна быть установлена в значение 1 перед импортом scipy и scikit-learn.
export SCIPY_ARRAY_API=1
Обратите внимание, что эта переменная среды предназначена для временного использования. Более подробную информацию можно найти в документации SciPy по Array API.
Некоторые оценщики Scikit-learn, которые в основном полагаются на NumPy (в отличие от использования Cython) для реализации алгоритмической логики своих методов fit, predict или transform, могут быть настроены на принятие входных данных любого совместимого с Array API типа данных и автоматически перенаправлять операции в соответствующее подпространство вместо NumPy.
На данном этапе эта поддержка считается экспериментальной и должна быть явно включена, как описано ниже.
Примечание
В настоящее время известно, что только array-api-strict, cupy, и PyTorch работают с оценщиками Scikit-learn.
В следующем видео представлен обзор принципов проектирования стандарта и того, как он способствует межпрограммной совместимости между библиотеками работы с массивами:
- Scikit-learn на GPU с Array API от Томаса Фэна на PyData NYC 2023.
11.1.1. Пример использования
Вот пример кода, демонстрирующий, как использовать CuPy для запуска LinearDiscriminantAnalysis на GPU:
>>> from sklearn.datasets import make_classification >>> from sklearn import config_context >>> from sklearn.discriminant_analysis import LinearDiscriminantAnalysis >>> import cupy >>> X_np, y_np = make_classification(random_state=0) >>> X_cu = cupy.asarray(X_np) >>> y_cu = cupy.asarray(y_np) >>> X_cu.device <CUDA Device 0> >>> with config_context(array_api_dispatch=True): ... lda = LinearDiscriminantAnalysis() ... X_trans = lda.fit_transform(X_cu, y_cu) >>> X_trans.device <CUDA Device 0>
После обучения модели атрибуты, которые являются массивами, также будут из того же подпространства Array API, что и обучающие данные. Например, если для обучения использовалось подпространство Array API CuPy, то атрибуты после обучения будут находиться на GPU. Мы предоставляем экспериментальную _estimator_with_converted_arrays утилиту, которая копирует атрибуты оценщика из Array API в ndarray:
>>> from sklearn.utils._array_api import _estimator_with_converted_arrays >>> cupy_to_ndarray = lambda array : array.get() >>> lda_np = _estimator_with_converted_arrays(lda, cupy_to_ndarray) >>> X_trans = lda_np.transform(X_np) >>> type(X_trans) <class 'numpy.ndarray'>
11.1.1.1. Поддержка PyTorch
Для поддержки тензоров PyTorch нужно установить array_api_dispatch=True и передать тензоры напрямую:
>>> import torch >>> X_torch = torch.asarray(X_np, device="cuda", dtype=torch.float32) >>> y_torch = torch.asarray(y_np, device="cuda", dtype=torch.float32) >>> with config_context(array_api_dispatch=True): ... lda = LinearDiscriminantAnalysis() ... X_trans = lda.fit_transform(X_torch, y_torch) >>> type(X_trans) <class 'torch.Tensor'> >>> X_trans.device.type 'cuda'
11.1.2. Поддержка Array API-совместимых входных данных
Оценщики и другие инструменты в scikit-learn, поддерживающие совместимые с Array API входные данные.
11.1.2.1. Оценщики
-
decomposition.PCA(сsvd_solver="full",svd_solver="randomized"иpower_iteration_normalizer="QR") -
linear_model.Ridge(сsolver="svd") -
discriminant_analysis.LinearDiscriminantAnalysis(сsolver="svd") preprocessing.KernelCentererpreprocessing.LabelEncoderpreprocessing.MaxAbsScalerpreprocessing.MinMaxScalerpreprocessing.Normalizer
11.1.2.2. Мета-оценщики
Мета-оценщики, принимающие входные данные Array API при условии, что базовая оценка также это делает:
11.1.2.3. Метрики
sklearn.metrics.cluster.entropysklearn.metrics.accuracy_scoresklearn.metrics.d2_tweedie_scoresklearn.metrics.f1_scoresklearn.metrics.max_errorsklearn.metrics.mean_absolute_errorsklearn.metrics.mean_absolute_percentage_errorsklearn.metrics.mean_gamma_deviance-
sklearn.metrics.mean_poisson_deviance(требуется включение поддержки Array API для SciPy) sklearn.metrics.mean_squared_errorsklearn.metrics.mean_squared_log_errorsklearn.metrics.mean_tweedie_deviancesklearn.metrics.multilabel_confusion_matrixsklearn.metrics.pairwise.additive_chi2_kernelsklearn.metrics.pairwise.chi2_kernelsklearn.metrics.pairwise.cosine_similaritysklearn.metrics.pairwise.cosine_distances-
sklearn.metrics.pairwise.euclidean_distances(см. Примечание по поддержке устройств для float64) sklearn.metrics.pairwise.linear_kernelsklearn.metrics.pairwise.paired_cosine_distancessklearn.metrics.pairwise.paired_euclidean_distancessklearn.metrics.pairwise.polynomial_kernel-
sklearn.metrics.pairwise.rbf_kernel(см. Примечание по поддержке устройств для float64) sklearn.metrics.pairwise.sigmoid_kernelsklearn.metrics.precision_recall_fscore_supportsklearn.metrics.r2_scoresklearn.metrics.root_mean_squared_errorsklearn.metrics.root_mean_squared_log_errorsklearn.metrics.zero_one_loss
11.1.2.4. Инструменты
Ожидается, что охват будет расти со временем. Следуйте за выделенной задачей на GitHub, чтобы отслеживать прогресс.
11.1.2.5. Тип возвращаемых значений и атрибутов подгонки
При вызове функций или методов с входными данными, совместимыми с Array API, принято возвращать массивы того же типа контейнера массивов и устройства, что и входные данные.
Аналогично, когда оценщик подгоняется с входными данными, совместимыми с Array API, атрибуты подгонки будут массивами из той же библиотеки, что и входные данные, и храниться на том же устройстве. Методы predict и transform впоследствии ожидают входные данные из той же библиотеки массивов и устройства, что и данные, переданные методу fit.
Однако обратите внимание, что функции оценки, возвращающие скалярные значения, возвращают скалярные значения Python (обычно экземпляр float) вместо скалярного значения массива.
11.1.3. Общие проверки оценщиков
Добавьте тег array_api_support в набор тегов оценщика, чтобы указать, что он поддерживает Array API. Это позволит включить специальные проверки в рамках общих тестов для проверки того, что результаты оценщиков одинаковы при использовании обычных входных данных NumPy и входных данных Array API.
Для запуска этих проверок необходимо установить array_api_compat в вашей тестовой среде. Для запуска полного набора проверок необходимо установить как PyTorch, так и CuPy и иметь графический процессор. Проверки, которые не могут быть выполнены или имеют недостающие зависимости, будут автоматически пропущены. Поэтому важно запускать тесты с флагом -v, чтобы увидеть, какие проверки пропущены:
pip install array-api-compat # and other libraries as needed pytest -k "array_api" -v
11.1.3.1. Замечание о поддержке устройств MPS
В macOS PyTorch может использовать Metal Performance Shaders (MPS) для доступа к ускорителям аппаратного обеспечения (например, к внутренней компоненте графического процессора чипов M1 или M2). Однако поддержка устройств MPS для PyTorch на момент написания неполная. Более подробную информацию см. в следующем вопросе на GitHub:
Для включения поддержки MPS в PyTorch установите переменную окружения PYTORCH_ENABLE_MPS_FALLBACK=1 перед запуском тестов:
PYTORCH_ENABLE_MPS_FALLBACK=1 pytest -k "array_api" -v
На момент написания все тесты scikit-learn должны пройти, однако скорость вычислений не обязательно лучше, чем при использовании устройства CPU.
11.1.3.2. Замечание о поддержке устройств для float64
Некоторые операции в scikit-learn автоматически выполняют операции с плавающей запятой с float64 точностью, чтобы предотвратить переполнение и обеспечить правильность (например, metrics.pairwise.euclidean_distances). Однако некоторые комбинации пространств массивов и устройств, такие как PyTorch on MPS (см. Замечание о поддержке устройств MPS), не поддерживают тип данных float64. В этих случаях scikit-learn вернется к использованию типа данных float32 вместо этого. Это может привести к другому поведению (обычно числовым нестабильным результатам) по сравнению с отсутствием диспетчеризации Array API или использованием устройства с поддержкой float64.
© 2007–2025 The scikit-learn developers
Licensed under the 3-clause BSD License.
https://scikit-learn.org/1.6/modules/array_api.html