From 30fdae4024ab657348c7e27ee638fd51bf298041 Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 25 Sep 2026 22:32:48 +0200 Subject: [PATCH] Added get_array_module --- CHANGELOG.md | 1 + docs/source/api.md | 9 +++++++++ src/cunumpy/__init__.py | 2 ++ src/cunumpy/__init__.pyi | 1 + src/cunumpy/xp.py | 17 +++++++++++++++++ tests/unit/test_cunumpy.py | 23 +++++++++++++++++++++++ 6 files changed, 53 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 26eaec4..5cc84f7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] ### Added +- `xp.get_array_module(array)`: Return the array-api-compat module (`numpy`/`cupy`) matching a given array's own backend, regardless of the process-wide active backend. Mirrors `cupy.get_array_module`, but works in Pyodide and returns array-api-compat modules for consistency with `xp.xp`. - Pyodide NumPy support documentation and CI that installs the built wheel in Pyodide's WebAssembly runtime and runs compiler-free tests for arrays, conversions, contexts, and Python kernels without CuPy or Pyccel imports. - `test-compiled` extra for native Pyccel tests. The `test` extra is now compiler-free; `dev` continues to include compiled-test dependencies. - `xp.same_backend(*arrays)`: Return `True` if all given arrays live on the same backend. diff --git a/docs/source/api.md b/docs/source/api.md index 178d6df..3452596 100644 --- a/docs/source/api.md +++ b/docs/source/api.md @@ -16,6 +16,15 @@ Converts an array to the currently active backend. ### `get_backend(array)` Returns the name of the backend (`"numpy"` or `"cupy"`) for the given array. +### `get_array_module(array)` +Returns the array-api-compat module (`numpy` or `cupy`) matching the given array's own backend — not necessarily the currently active global backend. Useful for writing functions that dispatch correctly on whatever array they receive, independent of `set_backend`/`use_backend`: + +```python +def norm(array): + xp_ = xp.get_array_module(array) + return xp_.sqrt(xp_.sum(array**2)) +``` + ### `is_gpu(array)` Returns `True` if the array is stored on a GPU (CuPy). diff --git a/src/cunumpy/__init__.py b/src/cunumpy/__init__.py index a35a539..53438c5 100644 --- a/src/cunumpy/__init__.py +++ b/src/cunumpy/__init__.py @@ -6,6 +6,7 @@ from .xp import ( assert_same_backend, cupy_available, + get_array_module, get_backend, is_cpu, is_gpu, @@ -30,6 +31,7 @@ "assert_same_backend", "cupy_available", "cupy_backend", + "get_array_module", "get_backend", "is_cpu", "is_gpu", diff --git a/src/cunumpy/__init__.pyi b/src/cunumpy/__init__.pyi index dc26342..88a27af 100644 --- a/src/cunumpy/__init__.pyi +++ b/src/cunumpy/__init__.pyi @@ -14,6 +14,7 @@ def to_numpy(array: Any) -> np.ndarray: ... def to_cupy(array: Any) -> Any: ... def to_cunumpy(array: Any) -> Any: ... def cupy_available() -> bool: ... +def get_array_module(array: Any) -> Any: ... def get_backend(array: Any) -> str: ... def is_gpu(array: Any) -> bool: ... def is_cpu(array: Any) -> bool: ... diff --git a/src/cunumpy/xp.py b/src/cunumpy/xp.py index 46daea4..e1a95cf 100644 --- a/src/cunumpy/xp.py +++ b/src/cunumpy/xp.py @@ -182,6 +182,23 @@ def get_backend(array: Any) -> BackendType: return "cupy" if array_api_compat.is_cupy_array(array) else "numpy" +def get_array_module(array: Any) -> ModuleType: + """Return the array-api-compat module matching `array`'s own backend. + + Unlike `xp.xp`, which reflects the process-wide active backend, this + dispatches on the array itself -- useful for writing functions that + operate correctly regardless of what `set_backend`/`use_backend` last + selected. Mirrors `cupy.get_array_module`, but also works in Pyodide + (where `cupy` cannot be imported) and returns array-api-compat modules + for standard-conformant behavior, consistent with `xp.xp`. + """ + if get_backend(array) == "cupy": + import array_api_compat.cupy as cp + + return cp + return np + + def is_gpu(array: Any) -> bool: """Check if the array is stored on a GPU (CuPy).""" return get_backend(array) == "cupy" diff --git a/tests/unit/test_cunumpy.py b/tests/unit/test_cunumpy.py index 38bf4d4..a6f9755 100644 --- a/tests/unit/test_cunumpy.py +++ b/tests/unit/test_cunumpy.py @@ -29,6 +29,29 @@ def test_get_backend_and_is_gpu_cpu(): assert xp.is_cpu(arr) is True +def test_get_array_module_numpy(): + arr = np.array([1, 2, 3]) + mod = xp.get_array_module(arr) + assert "numpy" in mod.__name__ + assert mod.asarray(arr) is not None + + +def test_get_array_module_matches_array_not_global_backend(): + if not xp.cupy_available(): + pytest.skip("CuPy not installed or not functional") + + a_cpu = np.array([1, 2, 3]) + a_gpu = xp.to_cupy(a_cpu) + + with xp.use_backend("cupy"): + # Global backend is cupy, but the array itself is on the host. + assert "numpy" in xp.get_array_module(a_cpu).__name__ + + with xp.use_backend("numpy"): + # Global backend is numpy, but the array itself is on the device. + assert "cupy" in xp.get_array_module(a_gpu).__name__ + + def test_same_backend_trivially_true_for_zero_or_one_array(): assert xp.same_backend() is True assert xp.same_backend(np.array([1, 2, 3])) is True