diff --git a/docs/release-notes/4315.feat.md b/docs/release-notes/4315.feat.md new file mode 100644 index 0000000000..e833b6358b --- /dev/null +++ b/docs/release-notes/4315.feat.md @@ -0,0 +1 @@ +The `use_rep` parameter of {func}`scanpy.pp.neighbors`, {func}`scanpy.tl.tsne`, and {func}`scanpy.tl.dendrogram` now accepts {mod}`anndata.acc` accessors such as `A.X`, `A.layers["scaled"]`, or `A.obsm["pca"]`, and resolves strings into them if {attr}`scanpy.settings.preset` is {attr}`~scanpy.Preset.ScanpyV2Preview` {smaller}`P Angerer` diff --git a/pyproject.toml b/pyproject.toml index 32801799f5..b045e21354 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -99,7 +99,7 @@ scrublet = [ "scikit-image>=0.25" ] # highly_variable_genes method 'seurat_v3' skmisc = [ "scikit-misc>=0.5.1" ] illico = [ "illico>=0.6" ] -scanpy2 = [ "anndata>=0.13.2", "hv-anndata>=0.0.3a5", "igraph>=0.10.8", "scanpy[illico]", "scikit-misc>=0.5.1" ] +scanpy2 = [ "anndata>=0.13.3", "hv-anndata>=0.0.3a5", "igraph>=0.10.8", "scanpy[illico]", "scikit-misc>=0.5.1" ] [dependency-groups] dev = [ diff --git a/src/scanpy/get/get.py b/src/scanpy/get/get.py index 80e38249b1..56aea122eb 100644 --- a/src/scanpy/get/get.py +++ b/src/scanpy/get/get.py @@ -2,7 +2,8 @@ from __future__ import annotations -from collections.abc import Collection +import json +from collections.abc import Collection, Sequence from importlib.util import find_spec from typing import TYPE_CHECKING, TypedDict, overload @@ -19,7 +20,7 @@ from collections.abc import Iterable from typing import Any, Literal, Unpack - from anndata.acc import Idx2D + from anndata.acc import Idx2D, RefAcc from .._compat import DaskArray @@ -485,6 +486,8 @@ class _Rep(TypedDict, total=False): type ArrAcc = GraphAcc | LayerAcc | MultiAcc +type RepAcc = LayerAcc | MultiAcc +"""Accessor usable as a representation (`use_rep`), i.e. a non-graph 2D array.""" @overload @@ -790,3 +793,70 @@ def _get_vec( ref = A.resolve(ref, vec=True) _ref_dim(ref, dim=dim) return adata[ref] + + +def _resolve_rep(rep: RefAcc | str) -> RepAcc: + """Resolve a `rep`resentation string into a `LayerAcc`/`MultiAcc` using `anndata.acc`.""" + if isinstance(rep, str): + from anndata.acc import A + + rep = A.resolve(rep, vec=False) + if isinstance(rep, MultiAcc) and rep.dim != "obs": + msg = ( + f"Representation must be aligned to `obs`, but {rep!r} is aligned to `var`" + ) + raise ValueError(msg) + if isinstance(rep, LayerAcc | MultiAcc): + return rep + msg = ( + "Representation must be a `LayerAcc` (e.g. `A.X`, `A.layers[...]`) or a " + f"`MultiAcc` (e.g. `A.obsm[...]`), was {rep!r}" + ) + raise TypeError(msg) + + +def _rep_to_json(rep: RepAcc | str | None) -> str | list[str] | None: + """Serialize a `rep`resentation for storage in `.uns`. + + v1 strings (`'X'` or an `.obsm` key) are stored unchanged, + accessors (and hence v2 strings) as `anndata.acc` JSON inside a 1-element list, + e.g. `A.obsm['pca']` as `['["obsm", "pca"]']`. + + TODO: Once AnnData can store a heterogeneous list, store that instead of a 1-element list. + See https://github.com/scverse/anndata/issues/1979 + """ + from scanpy import settings + + if rep is None or ( + isinstance(rep, str) and settings.preset is not Preset.ScanpyV2Preview + ): + return rep + from anndata.acc import A + + return [json.dumps(A.to_json(_resolve_rep(rep)))] + + +def _rep_from_json(rep: str | Sequence[str | int | None] | None) -> RepAcc | str | None: + """Parse a `rep`resentation stored by `_rep_to_json`.""" + from scanpy import settings + + if rep is None: + return rep + if not isinstance(rep, str): + from anndata.acc import A + + if ( + isinstance(rep, Sequence | np.ndarray) + and len(rep) == 1 + and isinstance(rep[0], str) + ): + # see `_rep_to_json` + rep: Sequence[str | int | None] = json.loads(rep[0]) + return _resolve_rep(A.from_json(rep, vec=False)) + if settings.preset is Preset.ScanpyV2Preview: + from anndata.acc import A + + # a plain string was stored under the v1 preset, + # so interpret it as one instead of as an `anndata.acc` spec + return A.X if rep == "X" else A.obsm[rep] + return rep diff --git a/src/scanpy/neighbors/__init__.py b/src/scanpy/neighbors/__init__.py index 13ddc0558f..4d79762aaa 100644 --- a/src/scanpy/neighbors/__init__.py +++ b/src/scanpy/neighbors/__init__.py @@ -23,6 +23,7 @@ from .._keys import _EmbeddingKeys, _existing_preset_keys from .._utils import NeighborsView, _doc_params, get_literal_vals from .._utils.random import _accepts_legacy_random_state, _LegacyRng +from ..get.get import _rep_to_json from . import _connectivity from ._common import ( _get_indices_distances_from_dense_matrix, @@ -44,6 +45,7 @@ from numpy.typing import NDArray from .._utils.random import RNGLike, SeedLike + from ..get.get import RepAcc from ._types import ( KnnTransformerLike, RPForestDict, @@ -64,7 +66,7 @@ def neighbors( # noqa: PLR0913 n_pcs: int | None = None, *, distances: np.ndarray | SpBase | None = None, - use_rep: str | None = None, + use_rep: RepAcc | str | None = None, knn: bool = True, method: _Method = "umap", transformer: KnnTransformerLike | _KnownTransformer | None = None, @@ -249,7 +251,7 @@ def neighbors( # noqa: PLR0913 metric=metric, **meta_random_state, **({} if not metric_kwds else dict(metric_kwds=metric_kwds)), - **({} if use_rep is None else dict(use_rep=use_rep)), + **({} if use_rep is None else dict(use_rep=_rep_to_json(use_rep))), **({} if n_pcs is None else dict(n_pcs=n_pcs)), ) @@ -536,7 +538,7 @@ def compute_neighbors( n_neighbors: int = 30, n_pcs: int | None = None, *, - use_rep: str | None = None, + use_rep: RepAcc | str | None = None, knn: bool = True, method: _Method | None = "umap", transformer: KnnTransformerLike | _KnownTransformer | None = None, @@ -564,7 +566,7 @@ def compute_neighbors( if `method` is not `None`, `.connectivities`. """ - from ..tools._utils import _choose_representation + from ..tools._utils import _choose_representation_compat start_neighbors = logg.debug("computing neighbors") if transformer is not None and not isinstance(transformer, str): @@ -590,7 +592,7 @@ def compute_neighbors( self._rp_forest = None self.n_neighbors = n_neighbors self.knn = knn - x = _choose_representation(self._adata, use_rep=use_rep, n_pcs=n_pcs) + x = _choose_representation_compat(self._adata, use_rep=use_rep, n_pcs=n_pcs) 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/neighbors/_doc.py b/src/scanpy/neighbors/_doc.py index ef454e296e..e7c5cf32fd 100644 --- a/src/scanpy/neighbors/_doc.py +++ b/src/scanpy/neighbors/_doc.py @@ -9,13 +9,19 @@ ``.obsp[.uns[neighbors_key]['connectivities_key']]`` for connectivities. """ -doc_use_rep = """\ -use_rep - Use the indicated representation. `'X'` or any key for `.obsm` is valid. +doc_use_rep = r"""use_rep + Use the indicated representation: + a :class:`~anndata.acc.LayerAcc` (e.g. `A.X`, `A.layers[...]`) or + :class:`~anndata.acc.MultiAcc` (e.g. `A.obsm[...]`, `A.varm[...]`). + A :class:`str` is :meth:`~anndata.acc.AdAcc.resolve`\ d to one of those + if :attr:`scanpy.settings.preset` is :attr:`~scanpy.Preset.ScanpyV2Preview`, + otherwise interpreted as `'X'` or a key of `.obsm`. + If `None`, the representation is chosen automatically: - For `.n_vars` < :attr:`~scanpy.settings.N_PCS` (default: 50), `.X` is used, otherwise 'X_pca' is used. - If 'X_pca' is not present, it’s computed with default parameters or `n_pcs` if present.\ -""" + For `.n_vars` < :attr:`~scanpy.settings.N_PCS` (default: 50), `.X` is used, otherwise the PCA + representation (`.obsm['X_pca']`, or `.obsm['pca']` if it was computed under + :attr:`~scanpy.Preset.ScanpyV2Preview`). + If it is not present, it’s computed with default parameters or `n_pcs` if present.""" doc_n_pcs = """\ n_pcs diff --git a/src/scanpy/neighbors/_types.py b/src/scanpy/neighbors/_types.py index e26a1153a0..080fb8bc5d 100644 --- a/src/scanpy/neighbors/_types.py +++ b/src/scanpy/neighbors/_types.py @@ -98,5 +98,5 @@ class NeighborsParams(TypedDict): metric: _Metric | _MetricFn | None random_state: NotRequired[_LegacyRandom] metric_kwds: NotRequired[Mapping[str, Any]] - use_rep: NotRequired[str] + use_rep: NotRequired[str | list[str]] # see `scanpy.get.get._rep_to_json` n_pcs: NotRequired[int] diff --git a/src/scanpy/tools/_dendrogram.py b/src/scanpy/tools/_dendrogram.py index 9f276f7052..cd22530d7c 100644 --- a/src/scanpy/tools/_dendrogram.py +++ b/src/scanpy/tools/_dendrogram.py @@ -9,8 +9,9 @@ from .. import logging as logg from .._utils import _doc_params, raise_not_implemented_error_if_backed_type +from ..get.get import _rep_to_json from ..neighbors._doc import doc_n_pcs, doc_use_rep -from ._utils import _choose_representation +from ._utils import _choose_representation_compat if TYPE_CHECKING: from collections.abc import Sequence @@ -18,6 +19,8 @@ from anndata import AnnData + from ..get.get import RepAcc + @_doc_params(n_pcs=doc_n_pcs, use_rep=doc_use_rep) def dendrogram( # noqa: PLR0913 @@ -25,7 +28,7 @@ def dendrogram( # noqa: PLR0913 groupby: str | Sequence[str], *, n_pcs: int | None = None, - use_rep: str | None = None, + use_rep: RepAcc | str | None = None, var_names: Sequence[str] | None = None, use_raw: bool | None = None, cor_method: str = "pearson", @@ -125,7 +128,7 @@ def dendrogram( # noqa: PLR0913 if var_names is None: rep_df = pd.DataFrame( - _choose_representation(adata, use_rep=use_rep, n_pcs=n_pcs) + _choose_representation_compat(adata, use_rep=use_rep, n_pcs=n_pcs) ) categorical = adata.obs[groupby[0]] if len(groupby) > 1: @@ -167,7 +170,7 @@ def dendrogram( # noqa: PLR0913 dat = dict( linkage=z_var, groupby=groupby, - use_rep=use_rep, + use_rep=_rep_to_json(use_rep), cor_method=cor_method, linkage_method=linkage_method, categories_ordered=dendro_info["ivl"], diff --git a/src/scanpy/tools/_ingest.py b/src/scanpy/tools/_ingest.py index a70700b629..31f8ed5061 100644 --- a/src/scanpy/tools/_ingest.py +++ b/src/scanpy/tools/_ingest.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextlib from collections.abc import MutableMapping from typing import TYPE_CHECKING @@ -21,7 +22,9 @@ from .._utils._doctests import doctest_skipif from .._utils.random import _legacy_random_state, _LegacyRng from ..get import _check_mask +from ..get.get import MultiAcc, _rep_from_json from ..neighbors import FlatTree +from ._utils import _choose_representation_compat if TYPE_CHECKING: from collections.abc import Generator, Iterable @@ -31,6 +34,7 @@ from umap import UMAP from .._keys import _EmbeddingKeys + from ..get.get import RepAcc from ..neighbors import RPForestDict @@ -225,7 +229,7 @@ class Ingest: _rng: np.random.Generator | None # neighbors _rep: np.ndarray - _use_rep: str + _use_rep: RepAcc | str _metric: str _metric_kwds: dict[str, object] _n_neighbors: int @@ -323,8 +327,10 @@ def _init_neighbors(self, adata: AnnData, neighbors_key: str | None) -> None: self._n_neighbors = neighbors["params"]["n_neighbors"] if "use_rep" in neighbors["params"]: - self._use_rep = neighbors["params"]["use_rep"] - self._rep = adata.X if self._use_rep == "X" else adata.obsm[self._use_rep] + self._use_rep = _rep_from_json(neighbors["params"]["use_rep"]) + self._rep = _choose_representation_compat( + adata, use_rep=self._use_rep, n_pcs=None + ) elif "n_pcs" in neighbors["params"]: self._use_rep = "X_pca" self._n_pcs = neighbors["params"]["n_pcs"] @@ -422,10 +428,11 @@ def _same_rep(self): adata = self._adata_new if self._n_pcs is not None: return self._pca(self._n_pcs) - if self._use_rep == "X": - return adata.X - if self._use_rep in adata.obsm: - return adata.obsm[self._use_rep] + # fall back to `.X` if the representation is missing in the new object + with contextlib.suppress(KeyError, ValueError): + return _choose_representation_compat( + adata, use_rep=self._use_rep, n_pcs=None + ) return adata.X def fit(self, adata_new: AnnData) -> None: @@ -558,11 +565,14 @@ def to_adata_joint( self._obsm[key], )) - if self._use_rep not in ("X_pca", "X"): - adata.obsm[self._use_rep] = np.vstack(( - self._adata_ref.obsm[self._use_rep], - self._obsm["rep"], - )) + pca_keys = _existing_preset_keys(self._adata_ref, "pca") + skip = {"X", pca_keys.obsm if pca_keys else "X_pca"} + match self._use_rep: + case MultiAcc(dim="obs", k=key) | str(key) if key not in skip: + adata.obsm[key] = np.vstack(( + self._adata_ref.obsm[key], + self._obsm["rep"], + )) if keys := _existing_preset_keys(self._adata_ref, "umap"): adata.uns[keys.uns] = self._adata_ref.uns[keys.uns] diff --git a/src/scanpy/tools/_tsne.py b/src/scanpy/tools/_tsne.py index d394369856..7da7597328 100644 --- a/src/scanpy/tools/_tsne.py +++ b/src/scanpy/tools/_tsne.py @@ -9,13 +9,15 @@ from .._settings import Default, settings from .._utils import _doc_params, raise_not_implemented_error_if_backed_type from .._utils.random import _accepts_legacy_random_state, _legacy_random_state +from ..get.get import _rep_to_json from ..neighbors._doc import doc_n_pcs, doc_use_rep -from ._utils import _choose_representation +from ._utils import _choose_representation_compat if TYPE_CHECKING: from anndata import AnnData from .._utils.random import RNGLike, SeedLike + from ..get.get import RepAcc @_accepts_legacy_random_state(0) @@ -25,7 +27,7 @@ def tsne( # noqa: PLR0913 n_pcs: int | None = None, *, n_components: int = 2, - use_rep: str | None = None, + use_rep: RepAcc | str | None = None, perplexity: float = 30, metric: str = "euclidean", early_exaggeration: float = 12, @@ -105,7 +107,7 @@ def tsne( # noqa: PLR0913 start = logg.info("computing tSNE") keys = _embedding_keys("tsne", key_added) adata = adata.copy() if copy else adata - x = _choose_representation(adata, use_rep=use_rep, n_pcs=n_pcs) + x = _choose_representation_compat(adata, use_rep=use_rep, n_pcs=n_pcs) raise_not_implemented_error_if_backed_type(x, "tsne") # params for sklearn n_jobs = settings.n_jobs if n_jobs is None else n_jobs @@ -156,7 +158,7 @@ def tsne( # noqa: PLR0913 learning_rate=learning_rate, n_jobs=n_jobs, metric=metric, - use_rep=use_rep, + use_rep=_rep_to_json(use_rep), n_components=n_components, ) adata.obsm[keys.obsm] = x_tsne # annotate samples with tSNE coordinates diff --git a/src/scanpy/tools/_umap.py b/src/scanpy/tools/_umap.py index 2a89b41269..7a7988afe6 100644 --- a/src/scanpy/tools/_umap.py +++ b/src/scanpy/tools/_umap.py @@ -15,7 +15,8 @@ _legacy_random_state, _LegacyRng, ) -from ._utils import _choose_representation, get_init_pos_from_paga +from ..get.get import _rep_from_json +from ._utils import _choose_representation_compat, get_init_pos_from_paga if TYPE_CHECKING: from typing import Literal @@ -182,9 +183,9 @@ def umap( # noqa: PLR0913 init_coords = check_array(init_coords, dtype=np.float32, accept_sparse=False) neigh_params = neighbors["params"] - x = _choose_representation( + x = _choose_representation_compat( adata, - use_rep=neigh_params.get("use_rep", None), + use_rep=_rep_from_json(neigh_params.get("use_rep", None)), n_pcs=neigh_params.get("n_pcs", None), silent=True, ) diff --git a/src/scanpy/tools/_utils.py b/src/scanpy/tools/_utils.py index ae7592c200..cacb79c71d 100644 --- a/src/scanpy/tools/_utils.py +++ b/src/scanpy/tools/_utils.py @@ -7,50 +7,87 @@ from .. import logging as logg from .._compat import warn from .._keys import _existing_preset_keys -from .._settings import settings +from .._settings import Preset, settings from .._utils import _choose_graph +from ..get.get import LayerAcc, MultiAcc, _get_arr, _resolve_rep if TYPE_CHECKING: from anndata import AnnData from numpy.typing import NDArray from .._compat import CSBase, CSRBase + from ..get.get import RepAcc def _choose_representation( adata: AnnData, *, - use_rep: str | None, + use_rep: RepAcc | str | None, n_pcs: int | None, silent: bool = False, ) -> np.ndarray | CSRBase: # TODO: what else? + """Get the representation to compute on, resolving strings using `anndata.acc`.""" + return _choose_representation_compat( + adata, + use_rep=None if use_rep is None else _resolve_rep(use_rep), + n_pcs=n_pcs, + silent=silent, + ) + + +def _choose_representation_compat( + adata: AnnData, + *, + use_rep: RepAcc | str | None, + n_pcs: int | None, + silent: bool = False, +) -> np.ndarray | CSRBase: # TODO: what else? + """Get the representation to compute on. + + Treats strings as `.obsm` keys (or `'X'`) instead of `anndata.acc` specs when the preset is v1. + """ verbosity = settings.verbosity if silent and settings.verbosity > 1: settings.verbosity = 1 + if use_rep is not None and ( + not isinstance(use_rep, str) or settings.preset is Preset.ScanpyV2Preview + ): + use_rep = _resolve_rep(use_rep) if use_rep is None and n_pcs == 0: # backwards compat for specifying `.X` use_rep = "X" - if use_rep is None: - x = _get_pca_or_small_x(adata, n_pcs) - elif use_rep in adata.obsm and n_pcs is not None: - if n_pcs > adata.obsm[use_rep].shape[1]: - msg = ( - f"{use_rep} does not have enough Dimensions. Provide a " - "Representation with equal or more dimensions than" - "`n_pcs` or lower `n_pcs` " - ) + match use_rep: + case None: + x = _get_pca_or_small_x(adata, n_pcs) + case LayerAcc(): + x = _get_arr(adata, use_rep) + case MultiAcc(): + x = _slice_n_pcs(_get_arr(adata, use_rep), n_pcs, use_rep) + case str() if use_rep in adata.obsm: + x = _slice_n_pcs(adata.obsm[use_rep], n_pcs, use_rep) + case "X": + x = adata.X + case _: + msg = f"Did not find {use_rep} in `.obsm.keys()`. You need to compute it first." raise ValueError(msg) - x = adata.obsm[use_rep][:, :n_pcs] - elif use_rep in adata.obsm and n_pcs is None: - x = adata.obsm[use_rep] - elif use_rep == "X": - x = adata.X - else: - msg = f"Did not find {use_rep} in `.obsm.keys()`. You need to compute it first." - raise ValueError(msg) settings.verbosity = verbosity # resetting verbosity return x +def _slice_n_pcs[A: np.ndarray | CSBase]( + x: A, n_pcs: int | None, use_rep: RepAcc | str +) -> A: + if n_pcs is None: + return x + if n_pcs > x.shape[1]: + msg = ( + f"{use_rep} does not have enough Dimensions. Provide a " + "Representation with equal or more dimensions than" + "`n_pcs` or lower `n_pcs` " + ) + raise ValueError(msg) + return x[:, :n_pcs] + + def _get_pca_or_small_x(adata: AnnData, n_pcs: int | None) -> np.ndarray | CSRBase: from .._keys import _embedding_keys from ..preprocessing._pca import pca diff --git a/src/testing/scanpy/_pytest/marks.py b/src/testing/scanpy/_pytest/marks.py index 06bd51fb31..629d4e1766 100644 --- a/src/testing/scanpy/_pytest/marks.py +++ b/src/testing/scanpy/_pytest/marks.py @@ -62,7 +62,7 @@ def _generate_next_value_( req: Requirement scanpy2 = "scanpy[scanpy2]" - anndata_acc = "anndata>=0.13.0rc3" + anndata_acc = "anndata>=0.13.3" colour = "colour-science" dask = auto() diff --git a/tests/test_ingest.py b/tests/test_ingest.py index 6103946d1d..70f7a564c5 100644 --- a/tests/test_ingest.py +++ b/tests/test_ingest.py @@ -1,5 +1,7 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import anndata import numpy as np import pytest @@ -10,6 +12,10 @@ import scanpy as sc from scanpy import settings from testing.scanpy._helpers.data import pbmc68k_reduced +from testing.scanpy._pytest.marks import needs + +if TYPE_CHECKING or hasattr(anndata, "acc"): + from anndata.acc import A X = np.array( [ @@ -70,6 +76,20 @@ def test_representation(adatas): assert ing._obsm["rep"] is adata_new.X +@needs.anndata_acc +def test_representation_acc(adatas) -> None: + """An accessor `use_rep` round-trips through `.uns` and is used for the new data.""" + adata_ref, adata_new = (a.copy() for a in adatas) + adata_new.obsm["X_pca"] = adata_ref.obsm["X_pca"][: adata_new.n_obs] + sc.pp.neighbors(adata_ref, use_rep=A.obsm["X_pca"]) + + ing = sc.tl.Ingest(adata_ref) + ing.fit(adata_new) + + assert ing._use_rep == A.obsm["X_pca"] + np.testing.assert_array_equal(ing._obsm["rep"], adata_new.obsm["X_pca"]) + + @pytest.mark.parametrize("as_sparse", [False, True]) def test_pca_transform_uses_reference_mean( as_sparse, monkeypatch: pytest.MonkeyPatch diff --git a/tests/test_neighbors.py b/tests/test_neighbors.py index 158adfb65b..8ad17a6fc3 100644 --- a/tests/test_neighbors.py +++ b/tests/test_neighbors.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING +import anndata import numpy as np import pytest from anndata import AnnData @@ -11,13 +12,22 @@ import scanpy as sc from scanpy import Neighbors from scanpy._compat import CSBase +from scanpy.get.get import _rep_from_json from testing.scanpy._helpers.data import pbmc68k_reduced +from testing.scanpy._pytest.marks import needs if TYPE_CHECKING: + from collections.abc import Callable + from pathlib import Path from typing import Literal from pytest_mock import MockerFixture + from scanpy.get.get import RepAcc + +if TYPE_CHECKING or hasattr(anndata, "acc"): + from anndata.acc import A + # the input data X = [[1, 0], [3, 0], [5, 6], [0, 4]] @@ -248,19 +258,119 @@ def test_metrics_argument(): assert not np.allclose(no_knn_euclidean.distances, no_knn_manhattan.distances) -def test_use_rep_argument(): - rng = np.random.default_rng() - adata = AnnData(rng.standard_normal((30, 300))) +@pytest.fixture +def adata_pca() -> AnnData: + rng = np.random.default_rng(0) + adata = AnnData(rng.standard_normal((30, 300)).astype(np.float32)) sc.pp.pca(adata) - neigh_pca = Neighbors(adata) + return adata + + +def test_use_rep_argument(adata_pca: AnnData): + neigh_pca = Neighbors(adata_pca) neigh_pca.compute_neighbors(n_pcs=5, use_rep="X_pca") - neigh_none = Neighbors(adata) + neigh_none = Neighbors(adata_pca) neigh_none.compute_neighbors(n_pcs=5, use_rep=None) np.testing.assert_allclose( neigh_pca.distances.toarray(), neigh_none.distances.toarray() ) +@needs.anndata_acc +@pytest.mark.parametrize( + ("acc", "legacy"), + [ + pytest.param(lambda: A.X, "X", id="X"), + pytest.param(lambda: A.obsm["X_pca"], "X_pca", id="obsm"), + ], +) +def test_use_rep_acc( + adata_pca: AnnData, acc: Callable[[], RepAcc], legacy: str +) -> None: + """An accessor `use_rep` is equivalent to the legacy string spelling it.""" + expected = adata_pca.copy() + sc.pp.neighbors(expected, n_pcs=5, use_rep=legacy) + sc.pp.neighbors(adata_pca, n_pcs=5, use_rep=acc()) + np.testing.assert_allclose( + adata_pca.obsp["distances"].toarray(), expected.obsp["distances"].toarray() + ) + + +@needs.anndata_acc +def test_use_rep_acc_stored(adata_pca: AnnData, tmp_path: Path) -> None: + """A stored accessor `use_rep` survives a round trip and is understood by readers.""" + sc.pp.neighbors(adata_pca, use_rep=A.obsm["X_pca"]) + assert adata_pca.uns["neighbors"]["params"]["use_rep"] == ['["obsm", "X_pca"]'] + adata_pca.write_h5ad(path := tmp_path / "adata.h5ad") + adata = anndata.read_h5ad(path) + assert ( + _rep_from_json(adata.uns["neighbors"]["params"]["use_rep"]) == A.obsm["X_pca"] + ) + sc.tl.umap(adata) # reads `use_rep` back out of `.uns` + + +@needs.scanpy2 +def test_use_rep_spec(adata_pca: AnnData) -> None: + """Under the v2 preset, a `use_rep` string is an `anndata.acc` spec.""" + expected = adata_pca.copy() + sc.pp.neighbors(expected, use_rep=A.obsm["X_pca"]) + with sc.settings.override(preset=sc.Preset.ScanpyV2Preview): + sc.pp.neighbors(adata_pca, use_rep="obsm.X_pca") + np.testing.assert_allclose( + adata_pca.obsp["distances"].toarray(), expected.obsp["distances"].toarray() + ) + + +@needs.anndata_acc +@pytest.mark.parametrize( + ("use_rep", "preset", "exc_type", "match"), + [ + pytest.param( + lambda: A.varm["PCs"], + sc.Preset.ScanpyV1, + ValueError, + "aligned to `obs`", + id="not-obs-aligned", + ), + pytest.param( + lambda: "nonexistent", + sc.Preset.ScanpyV1, + ValueError, + "Did not find nonexistent", + id="missing", + ), + pytest.param( + lambda: "nonexistent", + sc.Preset.ScanpyV2Preview, + ValueError, + "Cannot parse accessor", + id="unparsable", + marks=needs.scanpy2, + ), + pytest.param( + lambda: 1.0, + sc.Preset.ScanpyV1, + TypeError, + "must be a `LayerAcc`", + id="wrong-type", + ), + ], +) +def test_use_rep_invalid( + adata_pca: AnnData, + use_rep: Callable[[], object], + preset: sc.Preset, + exc_type: type[Exception], + match: str, +) -> None: + """Invalid `use_rep` arguments are refused with a message saying why.""" + with ( + sc.settings.override(preset=preset), + pytest.raises(exc_type, match=match), + ): + sc.pp.neighbors(adata_pca, use_rep=use_rep()) # type: ignore[arg-type] + + @pytest.mark.parametrize("conv", [sparse.csr_matrix.toarray, sparse.csr_matrix]) # noqa: TID251 def test_restore_n_neighbors(neigh, conv): neigh.compute_neighbors(n_neighbors, method="gauss")