diff --git a/docs/conf.py b/docs/conf.py index f0e322babd..ec76638f7a 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -191,11 +191,11 @@ "pp.highly_variable_genes": (["np", "sp", "da"], ["da[sp[csc]]"]), "pp.log1p": (["np", "sp", "da"], []), "pp.neighbors": (["np", "sp"], []), - "pp.normalize_total": (["np", "sp[csr]", "da"], []), + "pp.normalize_total": (["np", "sp[csr]", "da", "aa"], []), "pp.pca": (["np", "sp", "da"], ["da[sp[csc]]"]), "pp.regress_out": (["np"], []), "pp.sample": (["np", "sp", "da"], []), - "pp.scale": (["np", "sp", "da"], []), + "pp.scale": (["np", "sp", "da", "aa"], []), "pp.scrublet": (["np", "sp"], []), "pp.scrublet_simulate_doublets": (["np", "sp"], []), "tl.dendrogram": (["np", "sp"], []), diff --git a/docs/extensions/array_support.py b/docs/extensions/array_support.py index b79499fb2a..ae6b183a22 100644 --- a/docs/extensions/array_support.py +++ b/docs/extensions/array_support.py @@ -54,11 +54,22 @@ def run(self) -> list[nodes.Node]: # noqa: D102 )) title = nodes.title("", "", *self.parse_inline(":ref:`array-support`")[0]) - rows = self._render_support_data(data) + rows = [ + *self._render_support_data(data), + self._render_row( + self._render_array_type(_docs.ArrayApi()), + support=_docs.ArrayApi() in array_types, + in_dask=False, + ), + ] return self._render_table(headers, rows, title=title) def _render_overview(self) -> list[nodes.Node]: - headers = ["Function", *(at.rst(short=True) for at in ALL_INNER)] + headers = [ + "Function", + *(at.rst(short=True) for at in ALL_INNER), + _docs.ArrayApi().rst(short=True), + ] rows: list[nodes.row] = [] for fn, (include, exclude) in self._array_support.items(): row_header, _ = self.parse_inline(f":func:`scanpy.{fn}`") @@ -71,6 +82,7 @@ def _render_overview(self) -> list[nodes.Node]: ALL_INNER, map(_docs.DaskArray, ALL_INNER), strict=True ) ), + self._render_support(_docs.ArrayApi() in ats), ] rows.append( nodes.row( diff --git a/docs/release-notes/4179.feat.md b/docs/release-notes/4179.feat.md new file mode 100644 index 0000000000..cf616609b1 --- /dev/null +++ b/docs/release-notes/4179.feat.md @@ -0,0 +1 @@ +Add Array-API support, enabling JAX and other array-api backends in `adata.X` {smaller}`A. Karesh` diff --git a/pyproject.toml b/pyproject.toml index dd2432f923..5047792c4f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,7 +53,8 @@ classifiers = [ dynamic = [ "version" ] dependencies = [ "anndata>=0.12.14", - "fast-array-utils[accel,sparse]>=1.4", + "array-api-compat", + "fast-array-utils[accel,sparse]>=1.5", "h5py>=3.11", "joblib", "matplotlib>=3.10", @@ -103,10 +104,12 @@ scanpy2 = [ "anndata>=0.13.3", "hv-anndata>=0.0.3a5", "igraph>=0.10.8", "scanpy[ [dependency-groups] dev = [ - "scipy-stubs", # static typing and IDE support - "towncrier", # release note management + "scipy-stubs", # static typing and IDE support + "towncrier", # release note management + "types-array-api", ] test = [ + "jax", # Array API tests "scanpy[dask-ml]", "scanpy[dask]", "scanpy[illico]", diff --git a/src/scanpy/_utils/__init__.py b/src/scanpy/_utils/__init__.py index ddc04fc2af..626bc2f0b1 100644 --- a/src/scanpy/_utils/__init__.py +++ b/src/scanpy/_utils/__init__.py @@ -30,6 +30,8 @@ import numpy as np import pandas as pd from anndata._core.sparse_dataset import BaseCompressedSparseDataset +from array_api_compat import array_namespace +from fast_array_utils.types import HasArrayNamespace from .. import logging as logg from .._compat import CSBase, DaskArray, SpBase, warn @@ -609,6 +611,20 @@ def axis_mul_or_truediv( *, allow_divide_by_zero: bool = True, out: ArrayLike | None = None, +) -> np.ndarray: + raise NotImplementedError + + +@axis_mul_or_truediv.register(np.ndarray) +def _( + x: np.ndarray, + /, + scaling_array: np.ndarray, + axis: Literal[0, 1], + op: Callable[[Any, Any], Any], + *, + allow_divide_by_zero: bool = True, + out: ArrayLike | None = None, ) -> np.ndarray: _check_op(op) scaling_array = _broadcast_axis(scaling_array, axis) @@ -731,8 +747,38 @@ def _[T: (DaskArray, np.ndarray)]( ) +@axis_mul_or_truediv.register(HasArrayNamespace) +def _( + x: HasArrayNamespace, + /, + scaling_array: np.ndarray, + axis: Literal[0, 1], + op: Callable[[Any, Any], Any], + *, + allow_divide_by_zero: bool = True, + out: ArrayLike | None = None, +) -> Any: + + _check_op(op) + scaling_array = _broadcast_axis(scaling_array, axis) + xp = array_namespace(x) + scaling_array = xp.asarray(scaling_array) + if op is mul: + return x * scaling_array + if not allow_divide_by_zero: + scaling_array = xp.where( + scaling_array == 0, xp.ones_like(scaling_array), scaling_array + ) + return x / scaling_array + + @singledispatch def axis_nnz(x: ArrayLike, /, axis: Literal[0, 1]) -> np.ndarray: + raise NotImplementedError + + +@axis_nnz.register(np.ndarray) +def _(x: np.ndarray, /, axis: Literal[0, 1]) -> np.ndarray: return np.count_nonzero(x, axis=axis) @@ -751,6 +797,12 @@ def _(x: DaskArray, /, axis: Literal[0, 1]) -> DaskArray: ) +@axis_nnz.register(HasArrayNamespace) +def _(x: HasArrayNamespace, /, axis: Literal[0, 1]) -> Any: + xp = array_namespace(x) + return xp.count_nonzero(x, axis=axis) + + @singledispatch def check_nonnegative_integers(x: _SupportedArray, /) -> bool | DaskArray: """Check values of X to ensure it is count data.""" @@ -772,6 +824,16 @@ def _check_nonnegative_integers_in_mem(x: _MemoryArray, /) -> bool: return not np.any((data % 1) != 0) +@check_nonnegative_integers.register(HasArrayNamespace) +def _check_nonnegative_integers_array_api(x: HasArrayNamespace, /) -> bool: + xp = array_namespace(x) + if bool(xp.any(x < 0)): + return False + if xp.isdtype(x.dtype, "integral"): + return True + return not bool(xp.any((x % 1) != 0)) + + @check_nonnegative_integers.register(DaskArray) def _check_nonnegative_integers_dask(x: DaskArray, /) -> DaskArray: return x.map_blocks(check_nonnegative_integers, dtype=bool, drop_axis=(0, 1)) diff --git a/src/scanpy/_utils/_docs.py b/src/scanpy/_utils/_docs.py index a14a32cd4a..1e7ced6c78 100644 --- a/src/scanpy/_utils/_docs.py +++ b/src/scanpy/_utils/_docs.py @@ -12,7 +12,7 @@ from typing import Literal -__all__ = ["ArrayType", "DaskArray", "Numpy", "ScipySparse", "parse"] +__all__ = ["ArrayApi", "ArrayType", "DaskArray", "Numpy", "ScipySparse", "parse"] class ArrayType(ABC): @@ -32,6 +32,16 @@ def rst(self, *, short: bool = False) -> str: # pragma: no cover return f":class:`{'~' if short else ''}{self}`" +@dataclass(unsafe_hash=True, frozen=True) +class ArrayApi(ArrayType): + def __str__(self) -> str: # pragma: no cover + return "array-api" + + def rst(self, *, short: bool = False) -> str: # pragma: no cover + # No single class to link to, so link to the standard itself + return "`Array API `__" + + @dataclass(unsafe_hash=True, frozen=True) class ScipySparse(ArrayType): format: Literal["csr", "csc"] @@ -79,7 +89,7 @@ def parse( yield from (t for t in parse(include) if t not in excluded) return - inner_includes = [i for i in include if not i.startswith("da")] + inner_includes = [i for i in include if not i.startswith(("da", "aa"))] for t in include: if ( match := re.fullmatch(r"([^\[]+)(?:\[(.+)\])?", t) @@ -103,6 +113,11 @@ def _parse_mod( msg = f"`np` takes no tags {tags!r}" raise ValueError(msg) yield Numpy() + case "aa": + if tags: # pragma: no cover + msg = f"`aa` takes no tags {tags!r}" + raise ValueError(msg) + yield ArrayApi() case "sp": if tags - {"csr", "csc"}: # pragma: no cover msg = f"invalid tags {tags!r}" diff --git a/src/scanpy/metrics/_common.py b/src/scanpy/metrics/_common.py index c9c90e1dc8..4cce330f73 100644 --- a/src/scanpy/metrics/_common.py +++ b/src/scanpy/metrics/_common.py @@ -7,13 +7,12 @@ import numpy as np import pandas as pd +from fast_array_utils.types import HasArrayNamespace from .._compat import CSRBase, DaskArray, SpBase, fullname, warn from .._utils import NeighborsView if TYPE_CHECKING: - from typing import NoReturn - from anndata import AnnData from numpy.typing import NDArray @@ -90,10 +89,12 @@ def _resolve_vals[T: NDArray | DaskArray](val: T) -> T: ... def _resolve_vals(val: SpBase) -> CSRBase: ... @overload def _resolve_vals(val: pd.DataFrame | pd.Series) -> NDArray: ... +@overload +def _resolve_vals(val: HasArrayNamespace) -> NDArray: ... @singledispatch -def _resolve_vals(val: object) -> NoReturn: +def _resolve_vals(val: object): msg = f"Unsupported type {type(val)}" raise TypeError(msg) @@ -122,6 +123,12 @@ def _(val: pd.DataFrame | pd.Series) -> NDArray: return val.to_numpy() +@_resolve_vals.register(HasArrayNamespace) +def _resolve_vals_array_api(val: HasArrayNamespace) -> NDArray: + # Moran's I / Geary's C use numba kernels, so convert at the boundary + return np.asarray(val) + + def _vals_heterogeneous[V: NDArray | CSRBase]( vals: V, ) -> tuple[V, NDArray[np.bool] | slice, NDArray[np.float64]]: diff --git a/src/scanpy/neighbors/__init__.py b/src/scanpy/neighbors/__init__.py index 3197a6b45f..8ee2add54f 100644 --- a/src/scanpy/neighbors/__init__.py +++ b/src/scanpy/neighbors/__init__.py @@ -13,6 +13,7 @@ import numpy as np import scipy +from fast_array_utils.types import HasArrayNamespace from packaging.version import Version from scipy import sparse @@ -592,7 +593,11 @@ def compute_neighbors( self._rp_forest = None self.n_neighbors = n_neighbors self.knn = knn + x = _choose_representation_compat(self._adata, use_rep=use_rep, n_pcs=n_pcs) + if isinstance(x, HasArrayNamespace): + # sklearn transformers require numpy, so need to convert at boundary + x = np.asarray(x) self._distances = transformer.fit_transform(x) knn_indices, knn_distances = _get_indices_distances_from_sparse_matrix( self._distances, n_neighbors diff --git a/src/scanpy/preprocessing/_highly_variable_genes.py b/src/scanpy/preprocessing/_highly_variable_genes.py index eb9207e7fe..ee83c4e390 100644 --- a/src/scanpy/preprocessing/_highly_variable_genes.py +++ b/src/scanpy/preprocessing/_highly_variable_genes.py @@ -10,7 +10,9 @@ import numpy as np import pandas as pd from anndata import AnnData +from array_api_compat import array_namespace from fast_array_utils import stats +from fast_array_utils.types import HasArrayNamespace from .. import logging as logg from .._compat import CSBase, CSRBase, DaskArray, warn @@ -408,12 +410,17 @@ def _highly_variable_genes_single_batch( # use out if possible. only possible since we copy the data matrix if isinstance(x, np.ndarray): np.expm1(x, out=x) + elif isinstance(x, HasArrayNamespace): + xp = array_namespace(x) + x = xp.expm1(x) else: x = np.expm1(x) mean, var = materialize_as_ndarray(stats.mean_var(x, axis=0, correction=1)) # now actually compute the dispersion - mean[mean == 0] = 1e-12 # set entries equal to zero to small value + # JAX arrays are immutable, so in-place assignment (mean[mean == 0] = ...) + # fails; np.where allocates a fresh array instead + mean = np.where(mean == 0, 1e-12, mean) # set zero entries to a small value dispersion = var / mean if flavor == "seurat": # logarithmized mean as in Seurat dispersion[dispersion == 0] = np.nan diff --git a/src/scanpy/preprocessing/_normalization.py b/src/scanpy/preprocessing/_normalization.py index d0ce3ac175..affadfa7b8 100644 --- a/src/scanpy/preprocessing/_normalization.py +++ b/src/scanpy/preprocessing/_normalization.py @@ -5,6 +5,7 @@ import numba import numpy as np +from array_api_compat import array_namespace from fast_array_utils import stats from fast_array_utils.numba import njit @@ -15,14 +16,19 @@ if TYPE_CHECKING: from anndata import AnnData + from fast_array_utils.types import HasArrayNamespace -def _compute_nnz_median(counts: np.ndarray | DaskArray) -> np.floating: +def _compute_nnz_median( + counts: np.ndarray | DaskArray | HasArrayNamespace, +) -> np.floating: """Given a 1D array of counts, compute the median of the non-zero counts.""" if isinstance(counts, DaskArray): counts = counts.compute() + + xp = array_namespace(counts) counts_greater_than_zero = counts[counts > 0] - median = np.median(counts_greater_than_zero) + median = xp.median(counts_greater_than_zero) return median diff --git a/src/scanpy/preprocessing/_scale.py b/src/scanpy/preprocessing/_scale.py index a94b18cc3b..bc5921588d 100644 --- a/src/scanpy/preprocessing/_scale.py +++ b/src/scanpy/preprocessing/_scale.py @@ -7,8 +7,10 @@ import numba import numpy as np from anndata import AnnData +from array_api_compat import array_namespace from fast_array_utils.numba import njit from fast_array_utils.stats import mean_var +from fast_array_utils.types import HasArrayNamespace from scverse_misc import Deprecation, deprecated_arg from .. import logging as logg @@ -32,12 +34,18 @@ from ..get.get import Mask type _Array = CSBase | np.ndarray | DaskArray +type _Stat = NDArray[np.float64] | DaskArray @singledispatch def clip[A: _Array]( x: ArrayLike | A, *, max_value: float, zero_center: bool = True ) -> A: + raise NotImplementedError + + +@clip.register(np.ndarray) +def _(x: np.ndarray, *, max_value: float, zero_center: bool = True) -> np.ndarray: return clip_array(x, max_value=max_value, zero_center=zero_center) @@ -54,6 +62,12 @@ def _(x: DaskArray, *, max_value: float, zero_center: bool = True) -> DaskArray: ) +@clip.register(HasArrayNamespace) +def _(x, *, max_value: float, zero_center: bool = True): + xp = array_namespace(x) + return xp.clip(x, min=-max_value if zero_center else None, max=max_value) + + @njit def clip_array( x: NDArray[np.floating], /, *, max_value: float, zero_center: bool @@ -74,6 +88,49 @@ def clip_array( return x +def _cast_to_float[A: _Array](x: A) -> A: + """Cast integer input to float, as scaling leads to float results.""" + msg = ( + "... as scaling leads to float results, integer " + "input is cast to float, returning copy." + ) + if isinstance(x, np.ndarray | CSBase | DaskArray): + if not np.issubdtype(x.dtype, np.integer): + return x + logg.info(msg) + return x.astype(np.float64) + xp = array_namespace(x) + if not xp.isdtype(x.dtype, "integral"): + return x + logg.info(msg) + return xp.astype(x, xp.float64) + + +def _center_and_std[A: _Array](x: A, *, zero_center: bool) -> tuple[A, _Stat, _Stat]: + """Subtract the mean (if `zero_center`) and return the standard deviation.""" + mean, var = mean_var(x, axis=0, correction=1) + + if isinstance(x, np.ndarray | CSBase | DaskArray): + std = np.sqrt(var) + std[std == 0] = 1 + if zero_center: + if isinstance(x, CSBase) or ( + isinstance(x, DaskArray) and isinstance(x._meta, CSBase) + ): + msg = "zero-centering a sparse array/matrix densifies it." + warn(msg, UserWarning) + x -= mean + x = dematrix(x) + else: + xp = array_namespace(x) + std = xp.sqrt(var) + std = xp.where(std == 0, xp.ones_like(std), std) + if zero_center: + x = x - mean + + return x, mean, std + + @_doc_params( mask=doc_mask( "Restrict both the derivation of scaling parameters and the scaling itself\n" @@ -169,6 +226,7 @@ def scale[A: _Array]( @scale.register(DaskArray) @scale.register(CSBase) # `stacklevel=2` skips `singledispatch`’s dispatcher frame +@scale.register(HasArrayNamespace) @deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def scale_array[A: _Array]( x: A, @@ -179,14 +237,7 @@ def scale_array[A: _Array]( return_mean_std: bool = False, mask: NDArray[np.bool] | None = None, mask_obs: NDArray[np.bool] | None = None, -) -> ( - A - | tuple[ - A, - NDArray[np.float64] | DaskArray, - NDArray[np.float64], - ] -): +) -> A | tuple[A, _Stat, _Stat]: if copy: x = x.copy() @@ -199,13 +250,7 @@ def scale_array[A: _Array]( logg.info( # Be careful of what? This should be more specific "... be careful when using `max_value` without `zero_center`." ) - - if np.issubdtype(x.dtype, np.integer): - logg.info( - "... as scaling leads to float results, integer " - "input is cast to float, returning copy." - ) - x = x.astype(np.float64) + x = _cast_to_float(x) mask = _mask_arg(mask, mask_obs, dim="obs") mask = ( @@ -224,17 +269,7 @@ def scale_array[A: _Array]( return_mean_std=return_mean_std, ) - mean, var = mean_var(x, axis=0, correction=1) - std = np.sqrt(var) - std[std == 0] = 1 - if zero_center: - if isinstance(x, CSBase) or ( - isinstance(x, DaskArray) and isinstance(x._meta, CSBase) - ): - msg = "zero-centering a sparse array/matrix densifies it." - warn(msg, UserWarning) - x -= mean - x = dematrix(x) + x, mean, std = _center_and_std(x, zero_center=zero_center) x = axis_mul_or_truediv( x, @@ -260,14 +295,7 @@ def scale_array_masked[A: _Array]( zero_center: bool = True, max_value: float | None = None, return_mean_std: bool = False, -) -> ( - A - | tuple[ - A, - NDArray[np.float64] | DaskArray, - NDArray[np.float64], - ] -): +) -> A | tuple[A, _Stat, _Stat]: if isinstance(x, CSBase) and not zero_center: if isinstance(x, CSCBase): x = x.tocsr() diff --git a/src/scanpy/preprocessing/_simple.py b/src/scanpy/preprocessing/_simple.py index 3d014d60a0..f00f68a150 100644 --- a/src/scanpy/preprocessing/_simple.py +++ b/src/scanpy/preprocessing/_simple.py @@ -15,9 +15,11 @@ import numba import numpy as np from anndata import AnnData +from array_api_compat import array_namespace from fast_array_utils import stats from fast_array_utils.conv import to_dense from fast_array_utils.numba import njit +from fast_array_utils.types import HasArrayNamespace from numpy._typing._array_like import NDArray from pandas.api.types import CategoricalDtype from sklearn.utils import check_array @@ -379,6 +381,15 @@ def log1p_array(x: np.ndarray, *, base: Number | None = None, copy: bool = False return x +@log1p.register(HasArrayNamespace) +def log1p_array_api(x, *, base: Number | None = None, copy: bool = False): + xp = array_namespace(x) + result = xp.log1p(x) + if base is not None: + result = result / float(np.log(base)) + return result + + @log1p.register(AnnData) def log1p_anndata( adata: AnnData, @@ -822,7 +833,9 @@ def sample( # noqa: PLR0912 return subset.to_memory() if data.isbacked else subset.copy() # overload 3: return array and indices - assert isinstance(subset, np.ndarray | CSBase | DaskArray), type(subset) + assert isinstance(subset, np.ndarray | CSBase | DaskArray | HasArrayNamespace), ( + type(subset) + ) if copy: subset = subset.copy() return subset, indices diff --git a/src/scanpy/tools/_rank_genes_groups.py b/src/scanpy/tools/_rank_genes_groups.py index 4c4ed63614..4595512a40 100644 --- a/src/scanpy/tools/_rank_genes_groups.py +++ b/src/scanpy/tools/_rank_genes_groups.py @@ -9,6 +9,7 @@ import pandas as pd from anndata import AnnData from fast_array_utils.numba import njit +from fast_array_utils.types import HasArrayNamespace from scipy import sparse from scverse_misc import Deprecation, deprecated_arg @@ -28,7 +29,6 @@ ) from ..get import _check_mask, _get_arr, aggregate from ..get._aggregated import _chan_combine -from ..get.get import _mask_arg if TYPE_CHECKING: from collections.abc import Generator, Iterable @@ -290,6 +290,8 @@ def __init__( adata_comp = adata.raw x = adata_comp.X raise_not_implemented_error_if_backed_type(x, "rank_genes_groups") + if isinstance(x, HasArrayNamespace) and not isinstance(x, np.ndarray): + x = np.asarray(x) # for correct getnnz calculation if isinstance(x, CSBase): @@ -886,7 +888,11 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 >>> sc.pl.rank_genes_groups(adata) """ - mask = _mask_arg(mask, mask_var, dim="var") + from scanpy import settings + + # rank_genes_groups uses numba kernels internally, so need convert at entry. + if isinstance(mask_var, Default): + mask_var = settings.preset.rank_genes_groups.mask_var if isinstance(mean_in_log_space, Default): mean_in_log_space = settings.preset.rank_genes_groups.mean_in_log_space # If scanpy presets are used for v2, use illico - prevents the presets from showing the `wilcoxon_illico` method and allows us to silently replace `wilcoxon`'s implementation. diff --git a/src/testing/scanpy/_pytest/__init__.py b/src/testing/scanpy/_pytest/__init__.py index e7d4ebb027..73e4266f02 100644 --- a/src/testing/scanpy/_pytest/__init__.py +++ b/src/testing/scanpy/_pytest/__init__.py @@ -22,6 +22,12 @@ from collections.abc import Generator, Iterable, Mapping from pathlib import Path +if find_spec("jax"): + import jax + + # JAX defaults to 32-bit dtypes; enable 64-bit so results match the numpy + # reference values in the tests. + jax.config.update("jax_enable_x64", True) # noqa: FBT003 MARK_RETRY_DOWNLOAD = pytest.mark.flaky( reruns=5, diff --git a/src/testing/scanpy/_pytest/marks.py b/src/testing/scanpy/_pytest/marks.py index 629d4e1766..a6f08c652a 100644 --- a/src/testing/scanpy/_pytest/marks.py +++ b/src/testing/scanpy/_pytest/marks.py @@ -69,6 +69,7 @@ def _generate_next_value_( dask_ml = auto() fa2 = auto() gprofiler = "gprofiler-official" + jax = auto() leidenalg = auto() louvain = auto() openpyxl = auto() diff --git a/src/testing/scanpy/_pytest/params.py b/src/testing/scanpy/_pytest/params.py index 5ab37496ea..7d56259525 100644 --- a/src/testing/scanpy/_pytest/params.py +++ b/src/testing/scanpy/_pytest/params.py @@ -12,6 +12,11 @@ from .._helpers import as_dense_dask_array, as_sparse_dask_matrix from .._pytest.marks import needs +try: + from anndata.tests.helpers import as_dense_jax_array +except ImportError: + as_dense_jax_array = None + if TYPE_CHECKING: from collections.abc import Callable, Iterable from typing import Any, Literal @@ -68,7 +73,14 @@ def wrapper(a: np.ndarray) -> DaskArray: tuple[Literal["mem", "dask"], Literal["dense", "sparse"]], tuple[ParameterSet, ...], ] = { - ("mem", "dense"): (pytest.param(asarray, id="numpy_ndarray"),), + ("mem", "dense"): ( + pytest.param(asarray, id="numpy_ndarray"), + *( + [pytest.param(as_dense_jax_array, marks=[needs.jax], id="jax_array")] + if as_dense_jax_array is not None + else [] + ), + ), ("mem", "sparse"): ( pytest.param(sparse.csr_matrix, id="scipy_csr_mat"), # noqa: TID251 pytest.param(sparse.csc_matrix, id="scipy_csc_mat"), # noqa: TID251 diff --git a/tests/test_aggregated.py b/tests/test_aggregated.py index 7f1ef58aaf..90d3990170 100644 --- a/tests/test_aggregated.py +++ b/tests/test_aggregated.py @@ -17,6 +17,11 @@ from testing.scanpy._helpers.data import pbmc3k_processed from testing.scanpy._pytest.marks import needs from testing.scanpy._pytest.params import ARRAY_TYPES as ARRAY_TYPES_ALL +from testing.scanpy._pytest.params import ( + ARRAY_TYPES_MEM, + as_dense_jax_array, + param_with, +) if TYPE_CHECKING: from collections.abc import Callable @@ -28,7 +33,14 @@ from scanpy._compat import CSRBase VALID_ARRAY_TYPES = [ - at + param_with( + at, + marks=[ + pytest.mark.xfail(reason="aggregate not implemented for array-api arrays") + ], + ) + if at.id == "jax_array" + else at for at in ARRAY_TYPES_ALL if at.id not in { @@ -765,15 +777,19 @@ def test_nan() -> None: assert adata_agg.obs["n_obs_aggregated"].tolist() == [1, 2, 1] -@pytest.mark.parametrize("array_type", VALID_ARRAY_TYPES) -def test_var_no_catastrophic_cancellation(array_type) -> None: +@pytest.mark.parametrize("array_type", ARRAY_TYPES_MEM) +def test_var_no_catastrophic_cancellation( + request: pytest.FixtureRequest, array_type +) -> None: # Values of the form `offset + tiny_noise` make the textbook two-pass # formula sum(x**2)/n - (sum(x)/n)**2 lose ~all precision: both terms are # ~n*offset**2 ≈ 1e19 in float64 (precision ~1e3) but their difference is # the variance ~1e-3, far below the rounding noise. Welford's online - # algorithm and Chan's parallel combine (per chunk in dask, and for the - # zero-block merge in sparse paths) avoid the subtraction entirely. - + # algorithm avoids the subtraction entirely. + if array_type is as_dense_jax_array: + request.applymarker( + pytest.mark.xfail(reason="aggregate not implemented for jax arrays") + ) n_per_group, n_features = 1000, 4 offset, std = 1e8, 1e-3 groups = ["a", "b"] diff --git a/tests/test_highly_variable_genes.py b/tests/test_highly_variable_genes.py index 3f320b6d33..d8a7049ac0 100644 --- a/tests/test_highly_variable_genes.py +++ b/tests/test_highly_variable_genes.py @@ -378,7 +378,6 @@ def test_pearson_residuals_batch( @pytest.mark.parametrize("array_type", ARRAY_TYPES) def test_compare_to_upstream( *, - request: pytest.FixtureRequest, flavor: Literal["seurat", "cell_ranger"], params: Any, ref_path: Path, diff --git a/tests/test_pca.py b/tests/test_pca.py index a1fedface0..ab375f825f 100644 --- a/tests/test_pca.py +++ b/tests/test_pca.py @@ -19,6 +19,7 @@ from scanpy.preprocessing._pca._dask import _cov_sparse_dask from testing.scanpy import _helpers from testing.scanpy._helpers.data import pbmc3k_normalized +from testing.scanpy._pytest import params from testing.scanpy._pytest.marks import needs from testing.scanpy._pytest.params import ARRAY_TYPES as ARRAY_TYPES_ALL from testing.scanpy._pytest.params import param_with @@ -154,9 +155,9 @@ def possible_solvers( svd_solvers = {"arpack", "covariance_eigh"} case (type() as dc, False) if issubclass(dc, CSBase): svd_solvers = {"arpack", "randomized"} - case (helpers.asarray, True): + case (helpers.asarray | params.as_dense_jax_array, True): svd_solvers = {"auto", "full", "arpack", "randomized", "covariance_eigh"} - case (helpers.asarray, False): + case (helpers.asarray | params.as_dense_jax_array, False): svd_solvers = {"arpack", "randomized"} case _: pytest.fail(f"Unknown {array_type=} ({zero_center=}) ({id=})") diff --git a/tests/test_preprocessing.py b/tests/test_preprocessing.py index 2708cb3836..9128e78541 100644 --- a/tests/test_preprocessing.py +++ b/tests/test_preprocessing.py @@ -22,7 +22,10 @@ maybe_dask_process_context, ) from testing.scanpy._helpers.data import pbmc3k, pbmc68k_reduced -from testing.scanpy._pytest.params import ARRAY_TYPES, ARRAY_TYPES_SPARSE +from testing.scanpy._pytest.params import ( + ARRAY_TYPES, + ARRAY_TYPES_SPARSE, +) if TYPE_CHECKING: from collections.abc import Callable @@ -609,68 +612,30 @@ def test_recipe_weinreb(): @pytest.mark.parametrize("array_type", ARRAY_TYPES) -@pytest.mark.parametrize( - ("max_cells", "max_counts", "min_cells", "min_counts"), - [ - (100, None, None, None), - (None, 100, None, None), - (None, None, 20, None), - (None, None, None, 20), - ], -) -def test_filter_genes(array_type, max_cells, max_counts, min_cells, min_counts): +@pytest.mark.parametrize("arg", ["max_cells", "max_counts", "min_cells", "min_counts"]) +def test_filter_genes(array_type, arg: str) -> None: + kw = {arg: 100 if arg.startswith("max") else 20} adata = pbmc68k_reduced() adata.X = adata.raw.X adata_casted = adata.copy() adata_casted.X = array_type(adata_casted.raw.X) - sc.pp.filter_genes( - adata, - max_cells=max_cells, - max_counts=max_counts, - min_cells=min_cells, - min_counts=min_counts, - ) - sc.pp.filter_genes( - adata_casted, - max_cells=max_cells, - max_counts=max_counts, - min_cells=min_cells, - min_counts=min_counts, - ) + sc.pp.filter_genes(adata, **kw) + sc.pp.filter_genes(adata_casted, **kw) adata_casted.X = conv.to_dense(adata_casted.X, to_cpu_memory=True) adata.X = conv.to_dense(adata.X) assert_allclose(adata_casted.X, adata.X, rtol=1e-5, atol=1e-5) @pytest.mark.parametrize("array_type", ARRAY_TYPES) -@pytest.mark.parametrize( - ("max_genes", "max_counts", "min_genes", "min_counts"), - [ - pytest.param(100, None, None, None, id="max_genes"), - pytest.param(None, 100, None, None, id="max_counts"), - pytest.param(None, None, 20, None, id="min_genes"), - pytest.param(None, None, None, 20, id="min_counts"), - ], -) -def test_filter_cells(array_type, max_genes, max_counts, min_genes, min_counts): +@pytest.mark.parametrize("arg", ["max_genes", "max_counts", "min_genes", "min_counts"]) +def test_filter_cells(array_type, arg: str) -> None: + kw = {arg: 100 if arg.startswith("max") else 20} adata = pbmc68k_reduced() adata.X = adata.raw.X adata_casted = adata.copy() adata_casted.X = array_type(adata_casted.raw.X) - sc.pp.filter_cells( - adata, - max_genes=max_genes, - max_counts=max_counts, - min_genes=min_genes, - min_counts=min_counts, - ) - sc.pp.filter_cells( - adata_casted, - max_genes=max_genes, - max_counts=max_counts, - min_genes=min_genes, - min_counts=min_counts, - ) + sc.pp.filter_cells(adata, **kw) + sc.pp.filter_cells(adata_casted, **kw) adata_casted.X = conv.to_dense(adata_casted.X, to_cpu_memory=True) adata.X = conv.to_dense(adata.X) assert_allclose(adata_casted.X, adata.X, rtol=1e-5, atol=1e-5) diff --git a/tests/test_rank_genes_groups.py b/tests/test_rank_genes_groups.py index dc430f19e3..1fb32e1ab7 100644 --- a/tests/test_rank_genes_groups.py +++ b/tests/test_rank_genes_groups.py @@ -9,6 +9,7 @@ import pandas as pd import pytest from anndata import AnnData +from anndata.tests.helpers import asarray from scipy.stats import mannwhitneyu import scanpy as sc @@ -21,7 +22,10 @@ from testing.scanpy._helpers import random_mask from testing.scanpy._helpers.data import pbmc68k_reduced from testing.scanpy._pytest.marks import needs -from testing.scanpy._pytest.params import ARRAY_TYPES, ARRAY_TYPES_MEM +from testing.scanpy._pytest.params import ( + ARRAY_TYPES, + ARRAY_TYPES_MEM, +) if TYPE_CHECKING: from collections.abc import Callable, Sequence @@ -140,9 +144,9 @@ def test_results_layers( ) -> None: adata = get_example_data(array_type, rng=_LegacyRng(1234)) adata.layers["to_test"] = adata.X.copy() - x = adata.X.tolil() if isinstance(adata.X, CSBase) else adata.X - mask = np.random.default_rng().integers(0, 2, adata.shape, dtype=bool) - x[mask] = 0 + # zero out random entries in a writable numpy copy (jax arrays are immutable) + x = asarray(adata.X).copy() + x[np.random.default_rng().integers(0, 2, adata.shape, dtype=bool)] = 0 adata.X = array_type(x) scores = get_true_scores(data_dir, method)["scores"]