You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.
Dismiss alert
<metaproperty="og:description" content="The array API specification defines a standard API for all array manipulation libraries with a NumPy-like API. Some scikit-learn estimators that primarily rely on NumPy (as opposed to using Cython)..." />
<metaname="description" content="The array API specification defines a standard API for all array manipulation libraries with a NumPy-like API. Some scikit-learn estimators that primarily rely on NumPy (as opposed to using Cython)..." />
<title>13.2. Array API support (experimental) — scikit-learn 1.10.dev0 documentation</title>
<liclass="toctree-l1 has-children"><aclass="reference internal" href="../model_selection.html">3. Model selection and evaluation</a><details><summary><spanclass="toctree-toggle" role="presentation"><iclass="fa-solid fa-chevron-down"></i></span></summary><ul>
<liclass="toctree-l2"><aclass="reference internal" href="grid_search.html">3.2. Tuning the hyper-parameters of an estimator</a></li>
<liclass="toctree-l2"><aclass="reference internal" href="classification_threshold.html">3.3. Tuning the decision threshold for class prediction</a></li>
<liclass="toctree-l2"><aclass="reference internal" href="model_evaluation.html">3.4. Metrics and scoring: quantifying the quality of predictions</a></li>
<liclass="toctree-l2"><aclass="reference internal" href="learning_curve.html">3.5. Validation curves: plotting scores to evaluate models</a></li>
<liclass="toctree-l2"><aclass="reference internal" href="../computing/parallelism.html">10.3. Parallelism and resource management</a></li>
</ul>
</details></li>
<liclass="toctree-l1"><aclass="reference internal" href="../model_persistence.html">11. Model persistence</a></li>
<liclass="toctree-l1"><aclass="reference internal" href="../common_pitfalls.html">12. Common pitfalls and recommended practices</a></li>
<liclass="toctree-l1 current active has-children"><aclass="reference internal" href="../data_interoperability.html">13. Data Interoperability</a><detailsopen="open"><summary><spanclass="toctree-toggle" role="presentation"><iclass="fa-solid fa-chevron-down"></i></span></summary><ulclass="current">
<liclass="toctree-l2"><aclass="reference internal" href="df_output_transform.html">13.1. Pandas/Polars Output for Transformers with <codeclass="docutils literal notranslate"><spanclass="pre">set_output</span></code> API</a></li>
<liclass="toctree-l2 current active"><aclass="current reference internal" href="#">13.2. Array API support (experimental)</a></li>
</ul>
</details></li>
<liclass="toctree-l1"><aclass="reference internal" href="../machine_learning_map.html">14. Choosing the right estimator</a></li>
<liclass="toctree-l1"><aclass="reference internal" href="../presentations.html">15. External Resources, Videos and Talks</a></li>
<liclass="breadcrumb-item active" aria-current="page"><spanclass="ellipsis"><spanclass="section-number">13.2. </span>Array API support (experimental)</span></li>
</ul>
</nav>
</div>
</div>
</div>
</div>
<divid="searchbox"></div>
<articleclass="bd-article">
<sectionid="array-api-support-experimental">
<spanid="array-api"></span><h1><spanclass="section-number">13.2. </span>Array API support (experimental)<aclass="headerlink" href="#array-api-support-experimental" title="Link to this heading">#</a></h1>
a standard API for all array manipulation libraries with a NumPy-like API.</p>
<p>Some scikit-learn estimators that primarily rely on NumPy (as opposed to using
Cython) to implement the algorithmic logic of their <codeclass="docutils literal notranslate"><spanclass="pre">fit</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">predict</span></code> or
<codeclass="docutils literal notranslate"><spanclass="pre">transform</span></code> methods can be configured to accept any Array API compatible input
data structures and automatically dispatch operations to the underlying namespace
instead of relying on NumPy.</p>
<p>At this stage, this support is <strong>considered experimental</strong>, must be enabled
explicitly by the <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch</span></code> configuration and assumes that the latest
versions libraries are installed. See below for details.</p>
<p>The following video provides an overview of the standard’s design principles
and how it facilitates interoperability between array libraries:</p>
<ulclass="simple">
<li><p><aclass="reference external" href="https://www.youtube.com/watch?v=c_s8tr1AizA">Scikit-learn on GPUs with Array API</a>
by <aclass="reference external" href="https://github.com/thomasjpfan">Thomas Fan</a> at PyData NYC 2023.</p></li>
</ul>
<sectionid="supported-array-libraries">
<h2><spanclass="section-number">13.2.1. </span>Supported array libraries<aclass="headerlink" href="#supported-array-libraries" title="Link to this heading">#</a></h2>
<p>The following table lists the libraries and hardware for which we run automated
compliance tests on a regular basis. Other array API conforming libraries and
<td><p>See install link for driver setup; see <aclass="reference internal" href="#xpu-support"><spanclass="std std-ref">Note on Intel GPU support</span></a>;
see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a></p></td>
</tr>
</tbody>
</table>
</div>
<p>Coverage is expected to grow over time.</p>
</section>
<sectionid="enabling-array-api-support">
<h2><spanclass="section-number">13.2.2. </span>Enabling array API support<aclass="headerlink" href="#enabling-array-api-support" title="Link to this heading">#</a></h2>
<p>The configuration parameter <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch</span></code> needs to be set to <codeclass="docutils literal notranslate"><spanclass="pre">True</span></code> to enable array
API support. We recommend setting this configuration globally to ensure consistent
behaviour and prevent accidental mixing of array namespaces.
Note that in the examples below, we use a context manager (<aclass="reference internal" href="generated/sklearn.config_context.html#sklearn.config_context" title="sklearn.config_context"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">config_context</span></code></a>)
to avoid having to reset it to <codeclass="docutils literal notranslate"><spanclass="pre">False</span></code> at the end of every code snippet, so as to
not affect the rest of the documentation.</p>
<p>Scikit-learn’s support for the array API standard requires the environment variable
<codeclass="docutils literal notranslate"><spanclass="pre">SCIPY_ARRAY_API</span></code> to be set to <codeclass="docutils literal notranslate"><spanclass="pre">1</span></code> before importing <codeclass="docutils literal notranslate"><spanclass="pre">scipy</span></code> and <codeclass="docutils literal notranslate"><spanclass="pre">scikit-learn</span></code>:</p>
<p>Please note that this environment variable is intended for temporary use.
For more details, refer to SciPy’s <aclass="reference external" href="https://docs.scipy.org/doc/scipy/dev/api-dev/array_api.html#using-array-api-standard-support">Array API documentation</a>.</p>
<p>The array API functionality assumes that the latest versions of scikit-learn’s dependencies are
installed. Older versions might work, but we make no promises. While array API support is marked
as experimental, backwards compatibility is not guaranteed. In particular, when a newer version
of a dependency fixes a bug we will not introduce additional code to backport the fix or
maintain compatibility with older versions.</p>
<p>Scikit-learn accepts <aclass="reference internal" href="../glossary.html#term-array-like"><spanclass="xref std std-term">array-like</span></a> inputs for all <codeclass="xref py py-mod docutils literal notranslate"><spanclass="pre">metrics</span></code>
and some estimators. When <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch=False</span></code>, these inputs are converted
While this will successfully convert some array API inputs (e.g., JAX array),
we generally recommend setting <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch=True</span></code> when using array API inputs.
This is because NumPy conversion can often fail, e.g., torch tensor allocated on GPU.
Also note that the output array will be NumPy when <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch=False</span></code> whereas
when <codeclass="docutils literal notranslate"><spanclass="pre">array_api_dispatch=True</span></code> the output array will depend on the input array (see
<aclass="reference internal" href="#input-output-array-api"><spanclass="std std-ref">Input and output array type handling</span></a> for details).</p>
</section>
<sectionid="example-usage">
<h2><spanclass="section-number">13.2.3. </span>Example usage<aclass="headerlink" href="#example-usage" title="Link to this heading">#</a></h2>
<p>The example code snippet below demonstrates how to use <aclass="reference external" href="https://pytorch.org/">PyTorch</a> to run
<aclass="reference internal" href="generated/sklearn.discriminant_analysis.LinearDiscriminantAnalysis.html#sklearn.discriminant_analysis.LinearDiscriminantAnalysis" title="sklearn.discriminant_analysis.LinearDiscriminantAnalysis"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">LinearDiscriminantAnalysis</span></code></a> on a CUDA GPU:</p>
<p>This pattern works identically with any supported array library. For example,
replace <codeclass="docutils literal notranslate"><spanclass="pre">torch.asarray(...,</span><spanclass="pre">device="cuda")</span></code> with <codeclass="docutils literal notranslate"><spanclass="pre">cupy.asarray(...)</span></code> for CuPy
or <codeclass="docutils literal notranslate"><spanclass="pre">dpnp.asarray(...)</span></code> for dpnp. You can also target different devices within
PyTorch by changing the <codeclass="docutils literal notranslate"><spanclass="pre">device=</span></code> argument (e.g., <codeclass="docutils literal notranslate"><spanclass="pre">"cpu"</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">"xpu"</span></code>, <codeclass="docutils literal notranslate"><spanclass="pre">"mps"</span></code>).</p>
<p>After the model is trained, fitted attributes that are arrays will also be from
the same Array API namespace as the training data. For example, if PyTorch’s
CUDA namespace was used for training, then fitted attributes will be on the GPU.
Passing data in a different namespace or in a different device within the same
namespace to <codeclass="docutils literal notranslate"><spanclass="pre">transform</span></code> or <codeclass="docutils literal notranslate"><spanclass="pre">predict</span></code> is an error:</p>
<spanclass="gr">ValueError</span>: <spanclass="n">Inputs passed to LinearDiscriminantAnalysis.transform() must use the same namespace and the same device as those passed to fit()...</span>
</pre></div>
</div>
<sectionid="moving-estimators-between-devices">
<h3><spanclass="section-number">13.2.3.1. </span>Moving estimators between devices<aclass="headerlink" href="#moving-estimators-between-devices" title="Link to this heading">#</a></h3>
<p>We provide <codeclass="docutils literal notranslate"><spanclass="pre">move_estimator_to</span></code> to transfer an estimator’s array attributes
<spanid="array-api-supported"></span><h2><spanclass="section-number">13.2.4. </span>Support for array API compatible inputs<aclass="headerlink" href="#support-for-array-api-compatible-inputs" title="Link to this heading">#</a></h2>
<p>Estimators and other tools in scikit-learn that support array API compatible inputs.</p>
<sectionid="estimators">
<h3><spanclass="section-number">13.2.4.1. </span>Estimators<aclass="headerlink" href="#estimators" title="Link to this heading">#</a></h3>
<ulclass="simple">
<li><p><aclass="reference internal" href="generated/sklearn.covariance.LedoitWolf.html#sklearn.covariance.LedoitWolf" title="sklearn.covariance.LedoitWolf"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">covariance.LedoitWolf</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.covariance.OAS.html#sklearn.covariance.OAS" title="sklearn.covariance.OAS"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">covariance.OAS</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.linear_model.RidgeCV.html#sklearn.linear_model.RidgeCV" title="sklearn.linear_model.RidgeCV"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">linear_model.RidgeCV</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.linear_model.RidgeClassifierCV.html#sklearn.linear_model.RidgeClassifierCV" title="sklearn.linear_model.RidgeClassifierCV"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">linear_model.RidgeClassifierCV</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.preprocessing.StandardScaler.html#sklearn.preprocessing.StandardScaler" title="sklearn.preprocessing.StandardScaler"><codeclass="xref py py-class docutils literal notranslate"><spanclass="pre">preprocessing.StandardScaler</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.metrics.matthews_corrcoef.html#sklearn.metrics.matthews_corrcoef" title="sklearn.metrics.matthews_corrcoef"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">sklearn.metrics.matthews_corrcoef</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.metrics.pairwise.euclidean_distances.html#sklearn.metrics.pairwise.euclidean_distances" title="sklearn.metrics.pairwise.euclidean_distances"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">sklearn.metrics.pairwise.euclidean_distances</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<li><p><aclass="reference internal" href="generated/sklearn.metrics.pairwise.rbf_kernel.html#sklearn.metrics.pairwise.rbf_kernel" title="sklearn.metrics.pairwise.rbf_kernel"><codeclass="xref py py-func docutils literal notranslate"><spanclass="pre">sklearn.metrics.pairwise.rbf_kernel</span></code></a> (see <aclass="reference internal" href="#device-support-for-float64"><spanclass="std std-ref">Note on device support for float64</span></a>)</p></li>
<spanid="input-output-array-api"></span><h2><spanclass="section-number">13.2.5. </span>Input and output array type handling<aclass="headerlink" href="#input-and-output-array-type-handling" title="Link to this heading">#</a></h2>