From 94f3357eaf0aac19719875ad22009774d1b0a8d5 Mon Sep 17 00:00:00 2001 From: "Philipp A." Date: Tue, 1 Sep 2026 17:38:19 +0200 Subject: [PATCH 1/7] feat: add `anndata.acc` support to `mask` params --- docs/release-notes/4331.feat.md | 2 + docs/tutorials/basics/clustering-2017.ipynb | 8 +- pyproject.toml | 2 +- src/scanpy/_docs.py | 29 ++++++- src/scanpy/_settings/presets.py | 6 +- src/scanpy/_utils/__init__.py | 9 ++- src/scanpy/experimental/pp/_normalization.py | 27 ++++--- src/scanpy/get/_aggregated.py | 8 +- src/scanpy/get/get.py | 62 ++++++++++++++- src/scanpy/preprocessing/_docs.py | 13 ++-- .../preprocessing/_highly_variable_genes.py | 4 +- src/scanpy/preprocessing/_pca/__init__.py | 41 +++++----- src/scanpy/preprocessing/_scale.py | 60 ++++++++++----- src/scanpy/preprocessing/_simple.py | 16 ++-- src/scanpy/tools/_ingest.py | 4 +- src/scanpy/tools/_rank_genes_groups.py | 33 +++++--- tests/test_ingest.py | 3 +- tests/test_normalization.py | 11 ++- tests/test_pca.py | 35 +++++---- tests/test_rank_genes_groups.py | 4 +- tests/test_scaling.py | 77 +++++++++++++++---- 21 files changed, 324 insertions(+), 130 deletions(-) create mode 100644 docs/release-notes/4331.feat.md diff --git a/docs/release-notes/4331.feat.md b/docs/release-notes/4331.feat.md new file mode 100644 index 0000000000..37d00243ac --- /dev/null +++ b/docs/release-notes/4331.feat.md @@ -0,0 +1,2 @@ +Rename `mask_var`/`mask_obs` to `mask` in {func}`~scanpy.pp.pca`, {func}`~scanpy.pp.scale`, {func}`~scanpy.tl.rank_genes_groups`, and {func}`~scanpy.experimental.pp.normalize_pearson_residuals_pca`. +The new parameter takes {mod}`anndata.acc` references, e.g. `A.var["highly_variable"]` {smaller}`P Angerer` diff --git a/docs/tutorials/basics/clustering-2017.ipynb b/docs/tutorials/basics/clustering-2017.ipynb index 9e4601e5e3..d65870772c 100644 --- a/docs/tutorials/basics/clustering-2017.ipynb +++ b/docs/tutorials/basics/clustering-2017.ipynb @@ -1303,7 +1303,7 @@ } ], "source": [ - "sc.tl.rank_genes_groups(adata, \"leiden\", mask_var=\"highly_variable\", method=\"t-test\")\n", + "sc.tl.rank_genes_groups(adata, \"leiden\", mask=\"var.highly_variable\", method=\"t-test\")\n", "sc.pl.rank_genes_groups(adata, n_genes=25, sharey=False)" ] }, @@ -1355,7 +1355,7 @@ } ], "source": [ - "sc.tl.rank_genes_groups(adata, \"leiden\", mask_var=\"highly_variable\", method=\"wilcoxon\")\n", + "sc.tl.rank_genes_groups(adata, \"leiden\", mask=\"var.highly_variable\", method=\"wilcoxon\")\n", "sc.pl.rank_genes_groups(adata, n_genes=25, sharey=False)" ] }, @@ -1415,7 +1415,7 @@ ], "source": [ "sc.tl.rank_genes_groups(\n", - " adata, \"leiden\", mask_var=\"highly_variable\", method=\"logreg\", max_iter=1000\n", + " adata, \"leiden\", mask=\"var.highly_variable\", method=\"logreg\", max_iter=1000\n", ")\n", "sc.pl.rank_genes_groups(adata, n_genes=25, sharey=False)" ] @@ -2119,7 +2119,7 @@ "sc.tl.rank_genes_groups(\n", " adata,\n", " \"leiden\",\n", - " mask_var=\"highly_variable\",\n", + " mask=\"var.highly_variable\",\n", " groups=[\"0\"],\n", " reference=\"1\",\n", " method=\"wilcoxon\",\n", diff --git a/pyproject.toml b/pyproject.toml index 3183521eac..ab57866b81 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -124,7 +124,7 @@ docs = [ "plotly", # TODO: remove necessity for being able to import doc-linked classes "scanpy[dask-ml,leiden,paga,plotting,scanpy2,scrublet]", - "scanpydoc>=0.16.1", + "scanpydoc>=0.17.4", "scverse-misc[sphinx]", "sphinx>=9.1", "sphinx-autodoc-typehints>=1.25.2", diff --git a/src/scanpy/_docs.py b/src/scanpy/_docs.py index 2d22cdd817..ea490d8693 100644 --- a/src/scanpy/_docs.py +++ b/src/scanpy/_docs.py @@ -2,7 +2,34 @@ from __future__ import annotations -__all__ = ["doc_rng"] +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from typing import Literal + +__all__ = ["doc_mask", "doc_ref_compat", "doc_rng"] + +doc_ref_compat = ( + "If :attr:`scanpy.settings.preset` is :attr:`~scanpy.Preset.ScanpyV2Preview`, " + ":class:`str`\\ s are :meth:`anndata.acc.AdAcc.resolve`\\ d " + "to :class:`~anndata.acc.AdRef`\\ s (e.g. `'obs.is_control'`), " + "otherwise interpreted as column names." +) + + +def doc_mask(desc: str, *, dim: Literal["obs", "var"], extra: str = "") -> str: + """Docs for a `mask` parameter and its deprecated `mask_{dim}` alias.""" + return f"""\ +mask + {desc} + Given by a boolean array or a reference to one, e.g. `A.{dim}['selected']`. + :class:`str`\\ s are :meth:`anndata.acc.AdAcc.resolve`\\ d, e.g. `'{dim}.selected'`. +{f" {extra}\n" if extra else ""}\ +mask_{dim} + Deprecated alias of `mask`, where a :class:`str` always refers to + a column of :attr:`~anndata.AnnData.{dim}`. +""" + doc_rng = """\ rng diff --git a/src/scanpy/_settings/presets.py b/src/scanpy/_settings/presets.py index 2e37647442..1e3fe62cb1 100644 --- a/src/scanpy/_settings/presets.py +++ b/src/scanpy/_settings/presets.py @@ -93,7 +93,7 @@ class BasicEmbeddingPreset(NamedTuple): class RankGenesGroupsPreset(NamedTuple): method: DETest - mask_var: str | None + mask: str | None mean_in_log_space: bool @@ -245,10 +245,10 @@ def rank_genes_groups() -> Mapping[Preset, RankGenesGroupsPreset]: """ return { Preset.ScanpyV1: RankGenesGroupsPreset( - method="t-test", mask_var=None, mean_in_log_space=True + method="t-test", mask=None, mean_in_log_space=True ), Preset.ScanpyV2Preview: RankGenesGroupsPreset( - method="wilcoxon", mask_var=None, mean_in_log_space=False + method="wilcoxon", mask=None, mean_in_log_space=False ), } diff --git a/src/scanpy/_utils/__init__.py b/src/scanpy/_utils/__init__.py index c08def5140..dc7a9a8faf 100644 --- a/src/scanpy/_utils/__init__.py +++ b/src/scanpy/_utils/__init__.py @@ -69,12 +69,12 @@ "check_use_raw", "compute_association_matrix_of_groups", "descend_classes_and_funcs", + "dim_acc", "ensure_igraph", "get_igraph_from_adjacency", "get_literal_vals", "indent", "is_backed_type", - "obs_acc", "raise_not_implemented_error_if_backed_type", "renamed_arg", "sanitize_anndata", @@ -988,14 +988,15 @@ def _resolve_axis( raise ValueError(msg) -def obs_acc(obs_col: str) -> str | AdRef: +def dim_acc(col: str, *, dim: Literal["obs", "var"] = "obs") -> str | AdRef: + """Get reference to the `col`umn of `adata.{dim}` the way the active preset expects it.""" from .._settings import Preset, settings if settings.preset is Preset.ScanpyV2Preview: from anndata.acc import A - return A.obs[obs_col] - return obs_col + return getattr(A, dim)[col] + return col def is_backed_type(x: object, /) -> bool: diff --git a/src/scanpy/experimental/pp/_normalization.py b/src/scanpy/experimental/pp/_normalization.py index 8aacd99c8b..a0f34e7f4d 100644 --- a/src/scanpy/experimental/pp/_normalization.py +++ b/src/scanpy/experimental/pp/_normalization.py @@ -8,13 +8,14 @@ import numpy as np from anndata import AnnData +from scverse_misc import Deprecation, deprecated_arg from ... import logging as logg from ... import settings from ..._compat import CSBase, warn from ..._keys import _embedding_keys from ..._settings import Default -from ..._utils import _doc_params, check_nonnegative_integers, view_to_actual +from ..._utils import _doc_params, check_nonnegative_integers, dim_acc, view_to_actual from ..._utils.random import _accepts_legacy_random_state from ...experimental._docs import ( doc_adata, @@ -26,6 +27,7 @@ doc_pca_chunk, ) from ...get import _check_mask, _get_arr, _set_obs_rep +from ...get.get import _mask_arg from ...preprocessing._docs import doc_mask_var from ...preprocessing._pca import pca @@ -34,6 +36,7 @@ from typing import Any from ..._utils.random import RNGLike, SeedLike + from ...get.get import Mask def _pearson_residuals( @@ -163,11 +166,12 @@ def normalize_pearson_residuals( adata=doc_adata, dist_params=doc_dist_params, pca_chunk=doc_pca_chunk, - mask_var=doc_mask_var, + mask=doc_mask_var, check_values=doc_check_values, inplace=doc_inplace, ) @_accepts_legacy_random_state(0) +@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead.")) def normalize_pearson_residuals_pca( adata: AnnData, *, @@ -176,11 +180,11 @@ def normalize_pearson_residuals_pca( n_comps: int | None = 50, rng: SeedLike | RNGLike | None = None, kwargs_pca: Mapping[str, Any] = frozendict({}), - mask_var: np.ndarray | str | Default | None = Default( - "adata.var.get('highly_variable')" - ), + mask: Mask | Default | None = Default("adata.var.get('highly_variable')"), check_values: bool = True, inplace: bool = True, + # deprecated + mask_var: Mask | None = None, ) -> AnnData | None: """Apply analytic Pearson residual normalization and PCA, based on :cite:t:`Lause2021`. @@ -196,7 +200,7 @@ def normalize_pearson_residuals_pca( {adata} {dist_params} {pca_chunk} - {mask_var} + {mask} {check_values} {inplace} @@ -227,9 +231,14 @@ def normalize_pearson_residuals_pca( """ key_added = kwargs_pca.get("key_added", settings.preset.pca.key_added) keys = _embedding_keys("pca", key_added) - if isinstance(mask_var, Default): - mask_var = "highly_variable" if "highly_variable" in adata.var else None - mask_var = _check_mask(adata, mask_var, "var") + mask = _mask_arg(mask, mask_var, dim="var") + if isinstance(mask, Default): + mask = ( + dim_acc("highly_variable", dim="var") + if "highly_variable" in adata.var + else None + ) + mask_var = _check_mask(adata, mask, "var") if mask_var is not None: adata_sub = adata[:, mask_var].copy() diff --git a/src/scanpy/get/_aggregated.py b/src/scanpy/get/_aggregated.py index 4a17068bc5..b58ae16d76 100644 --- a/src/scanpy/get/_aggregated.py +++ b/src/scanpy/get/_aggregated.py @@ -13,8 +13,9 @@ from sklearn.utils.sparsefuncs import csc_median_axis_0 from .._compat import CSBase, CSRBase, DaskArray, warn +from .._docs import doc_ref_compat from .._settings import Preset, settings -from .._utils import _resolve_axis, get_literal_vals +from .._utils import _doc_params, _resolve_axis, get_literal_vals from .._utils._doctests import doctest_needs from ._kernels import ( agg_sum_csc, @@ -238,6 +239,7 @@ def _normalize_by[I: Idx2D | int]( return by_list, dim +@_doc_params(ref=doc_ref_compat) @doctest_needs("anndata_acc") def aggregate( adata: AnnData, @@ -336,9 +338,7 @@ def aggregate( var: 'n_cells' layers: 'mean', 'count_nonzero' - .. [#ref] If :attr:`scanpy.settings.preset` is :attr:`~scanpy.Preset.ScanpyV2Preview`, - :class:`str`\ s are :meth:`anndata.acc.AdAcc.resolve`\ d to :class:`~anndata.acc.AdRef`\ s, - otherwise interpreted as :attr:`anndata.AnnData.obs` columns. + .. [#ref] {ref} """ if not isinstance(adata, AnnData): diff --git a/src/scanpy/get/get.py b/src/scanpy/get/get.py index 56aea122eb..4582ab94ef 100644 --- a/src/scanpy/get/get.py +++ b/src/scanpy/get/get.py @@ -13,14 +13,15 @@ from numpy.typing import NDArray from .._compat import CSBase -from .._settings import Preset +from .._settings import Default, Preset +from .._utils import dim_acc if TYPE_CHECKING: import sys from collections.abc import Iterable from typing import Any, Literal, Unpack - from anndata.acc import Idx2D, RefAcc + from anndata.acc import RefAcc from .._compat import DaskArray @@ -31,13 +32,18 @@ if TYPE_CHECKING or find_spec("anndata.acc"): - from anndata.acc import AdRef, GraphAcc, LayerAcc, MultiAcc + from anndata.acc import AdRef, GraphAcc, Idx2D, LayerAcc, MultiAcc else: AdRef = type("AdRef", (), dict(__module__="anndata.acc")) GraphAcc = type("GraphAcc", (), dict(__module__="anndata.acc")) + # https://github.com/tox-dev/sphinx-autodoc-typehints/issues/764 + type Idx2D = object LayerAcc = type("LayerAcc", (), dict(__module__="anndata.acc")) MultiAcc = type("MultiAcc", (), dict(__module__="anndata.acc")) +type Mask = NDArray[np.bool] | AdRef[Idx2D | int, AnnData] | str +"""A boolean array, or a reference to one (see `_check_mask`).""" + # -------------------------------------------------------------------------------- # Plotting data helpers # -------------------------------------------------------------------------------- @@ -607,6 +613,22 @@ def _set_obs_rep( raise AssertionError(msg) +def _mask_arg[M]( + mask: M | Default, legacy: M | None, *, dim: Literal["obs", "var"] +) -> M | Default: + """Merge the `mask` argument with its deprecated `mask_{dim}` predecessor.""" + if legacy is not None: + if mask is not None and not isinstance(mask, Default): + msg = f"Pass either `mask` or `mask_{dim}`, not both." + raise TypeError(msg) + return dim_acc(legacy, dim=dim) if isinstance(legacy, str) else legacy + if isinstance(mask, str): + from anndata.acc import A + + return A.resolve(mask, vec=True) + return mask + + def _check_mask[M: NDArray[np.bool] | NDArray[np.floating] | pd.Series | None]( data: AnnData | np.ndarray | CSBase | DaskArray, mask: str | AdRef[Idx2D | int, AnnData] | M, @@ -622,7 +644,7 @@ def _check_mask[M: NDArray[np.bool] | NDArray[np.floating] | pd.Series | None]( Annotated data matrix or numpy array. mask Mask (or probabilities if `allow_probabilities=True`). - Either an appropriatley sized array, or name of a column. + Either an appropriatley sized array, or a reference to one. dim The dimension being masked. allow_probabilities @@ -815,6 +837,38 @@ def _resolve_rep(rep: RefAcc | str) -> RepAcc: raise TypeError(msg) +def _ref_to_json[M: NDArray | None](ref: AdRef | str | M) -> str | list[str] | M: + """Serialize a vector reference for storage in `.uns`, see `_rep_to_json`. + + Arrays (and v1 strings) are stored unchanged. + """ + from scanpy import settings + + if not isinstance(ref, AdRef) and not ( + isinstance(ref, str) and settings.preset is Preset.ScanpyV2Preview + ): + return ref + from anndata.acc import A + + return [json.dumps(A.to_json(_resolve_ref(ref)))] + + +def _ref_from_json[M: NDArray | None]( + ref: str | Sequence[str | int | None] | M, +) -> AdRef | str | M: + """Parse a vector reference stored by `_ref_to_json`.""" + if ( + isinstance(ref, Sequence | np.ndarray) + and not isinstance(ref, str) + and len(ref) == 1 + and isinstance(ref[0], str) + ): + from anndata.acc import A + + return A.from_json(json.loads(ref[0]), vec=True) + return ref + + def _rep_to_json(rep: RepAcc | str | None) -> str | list[str] | None: """Serialize a `rep`resentation for storage in `.uns`. diff --git a/src/scanpy/preprocessing/_docs.py b/src/scanpy/preprocessing/_docs.py index 132284c104..8318a750df 100644 --- a/src/scanpy/preprocessing/_docs.py +++ b/src/scanpy/preprocessing/_docs.py @@ -2,6 +2,8 @@ from __future__ import annotations +from .._docs import doc_mask + doc_adata_basic = """\ adata Annotated data matrix.\ @@ -15,12 +17,11 @@ If True, use `adata.raw.X` for expression values instead of `adata.X`.\ """ -doc_mask_var = """\ -mask_var - To run only on a certain set of genes given by a boolean array - or a string referring to an array in :attr:`~anndata.AnnData.var`. - By default, uses `.var['highly_variable']` if available, else everything. -""" +doc_mask_var = doc_mask( + "To run only on a certain set of genes.", + dim="var", + extra="By default, uses `.var['highly_variable']` if available, else everything.", +) doc_obs_qc_args = """\ qc_vars diff --git a/src/scanpy/preprocessing/_highly_variable_genes.py b/src/scanpy/preprocessing/_highly_variable_genes.py index 0f8f77d996..d76e76cfd3 100644 --- a/src/scanpy/preprocessing/_highly_variable_genes.py +++ b/src/scanpy/preprocessing/_highly_variable_genes.py @@ -17,7 +17,7 @@ from .._settings import Default, Verbosity, settings from .._utils import ( check_nonnegative_integers, - obs_acc, + dim_acc, raise_if_dask_feature_axis_chunked, sanitize_anndata, ) @@ -181,7 +181,7 @@ def _highly_variable_genes_seurat_v3( # noqa: PLR0912, PLR0915 ) if batch_key is not None: aggregated_mean_var = aggregate( - adata_agg, by=obs_acc("__hvg_v3_batch_info__"), func=["mean", "var"] + adata_agg, by=dim_acc("__hvg_v3_batch_info__"), func=["mean", "var"] ) aggregated_mean_var.layers["mean"], aggregated_mean_var.layers["var"] = ( materialize_as_ndarray( diff --git a/src/scanpy/preprocessing/_pca/__init__.py b/src/scanpy/preprocessing/_pca/__init__.py index 9df89de43a..39ebd954a5 100644 --- a/src/scanpy/preprocessing/_pca/__init__.py +++ b/src/scanpy/preprocessing/_pca/__init__.py @@ -4,15 +4,17 @@ import numpy as np from anndata import AnnData +from scverse_misc import Deprecation, deprecated_arg from ... import logging as logg from ..._compat import CSBase, DaskArray, warn from ..._docs import doc_rng from ..._keys import _embedding_keys -from ..._settings import Default, Preset, settings -from ..._utils import _doc_params, get_literal_vals, is_backed_type +from ..._settings import Default, settings +from ..._utils import _doc_params, dim_acc, get_literal_vals, is_backed_type from ..._utils.random import _accepts_legacy_random_state, _legacy_random_state from ...get import _check_mask, _get_arr +from ...get.get import _mask_arg, _ref_to_json from .._docs import doc_mask_var from ._compat import _pca_compat_sparse @@ -23,9 +25,10 @@ import dask_ml.decomposition as dmld import sklearn.decomposition as skld - from numpy.typing import DTypeLike, NDArray + from numpy.typing import DTypeLike from ..._utils.random import RNGLike, SeedLike + from ...get.get import Mask type MethodDaskML = type[dmld.PCA | dmld.IncrementalPCA | dmld.TruncatedSVD] @@ -48,8 +51,9 @@ type SvdSolver = SvdSolvDaskML | SvdSolvSkearn | SvdSolvPCACustom -@_doc_params(mask_var=doc_mask_var, rng=doc_rng) +@_doc_params(mask=doc_mask_var, rng=doc_rng) @_accepts_legacy_random_state(0) +@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead.")) def pca( # noqa: PLR0912, PLR0913, PLR0915 data: AnnData | np.ndarray | CSBase, n_comps: int | None = None, @@ -62,12 +66,11 @@ def pca( # noqa: PLR0912, PLR0913, PLR0915 chunk_size: int | None = None, rng: SeedLike | RNGLike | None = None, return_info: bool = False, - mask_var: NDArray[np.bool] | str | Default | None = Default( - "adata.var.get('highly_variable')" - ), + mask: Mask | Default | None = Default("adata.var.get('highly_variable')"), dtype: DTypeLike = "float32", key_added: str | Default | None = Default(preset=("pca", "key_added")), copy: bool = False, + mask_var: Mask | None = None, ) -> AnnData | np.ndarray | CSBase | None: r"""Principal component analysis :cite:p:`Pedregosa2011`. @@ -158,7 +161,7 @@ def pca( # noqa: PLR0912, PLR0913, PLR0915 return_info Only relevant when not passing an :class:`~anndata.AnnData`: see “Returns”. - {mask_var} + {mask} dtype Numpy data type string to which to convert the result. key_added @@ -218,17 +221,17 @@ def pca( # noqa: PLR0912, PLR0913, PLR0915 else: adata = AnnData(data) - if isinstance(mask_var, Default): - if "highly_variable" not in adata.var: - mask_var = None - elif settings.preset is Preset.ScanpyV2Preview: - mask_var = "var.highly_variable" - else: - mask_var = "highly_variable" - elif mask_var is not None and obsm is not None: - msg = "Argument `mask_var` is incompatible with `obsm`." + mask = _mask_arg(mask, mask_var, dim="var") + if isinstance(mask, Default): + mask = ( + dim_acc("highly_variable", dim="var") + if "highly_variable" in adata.var + else None + ) + elif mask is not None and obsm is not None: + msg = "Argument `mask` is incompatible with `obsm`." raise ValueError(msg) - mask_var_param, mask_var = mask_var, _check_mask(adata, mask_var, "var") + mask_param, mask_var = mask, _check_mask(adata, mask, "var") adata_comp = adata[:, mask_var] if mask_var is not None else adata if n_comps is None: @@ -353,7 +356,7 @@ def pca( # noqa: PLR0912, PLR0913, PLR0915 adata.uns[keys.uns] = dict( params=dict( zero_center=zero_center, - mask_var=mask_var_param, + mask_var=_ref_to_json(mask_param), **(dict(layer=layer) if layer is not None else {}), **(dict(obsm=obsm) if obsm is not None else {}), ), diff --git a/src/scanpy/preprocessing/_scale.py b/src/scanpy/preprocessing/_scale.py index 990f2d394f..7647e075c6 100644 --- a/src/scanpy/preprocessing/_scale.py +++ b/src/scanpy/preprocessing/_scale.py @@ -9,11 +9,14 @@ from anndata import AnnData from fast_array_utils.numba import njit from fast_array_utils.stats import mean_var +from scverse_misc import Deprecation, deprecated_arg from .. import logging as logg from .._compat import CSBase, CSCBase, CSRBase, DaskArray, warn +from .._docs import doc_mask from .._settings import Default, settings from .._utils import ( + _doc_params, axis_mul_or_truediv, check_array_function_arguments, dematrix, @@ -21,10 +24,13 @@ view_to_actual, ) from ..get import _check_mask, _get_arr, _set_obs_rep +from ..get.get import AdRef, _mask_arg if TYPE_CHECKING: from numpy.typing import ArrayLike, NDArray + from ..get.get import Mask + type _Array = CSBase | np.ndarray | DaskArray @@ -68,7 +74,16 @@ def clip_array( return x +@_doc_params( + mask=doc_mask( + "Restrict both the derivation of scaling parameters and the scaling itself\n" + " to a certain set of observations.", + dim="obs", + extra="This will transform data from csc to csr format if `issparse(data)`.", + ) +) @singledispatch +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) def scale[A: _Array]( data: AnnData | A, *, @@ -77,7 +92,9 @@ def scale[A: _Array]( copy: bool = False, layer: str | None = None, obsm: str | None = None, - mask_obs: NDArray[np.bool] | str | None = None, + mask: Mask | None = None, + # deprecated + mask_obs: Mask | None = None, ) -> AnnData | A | None: """Scale data to unit variance and zero mean. @@ -114,11 +131,7 @@ def scale[A: _Array]( If provided, which element of layers to scale. obsm If provided, which element of obsm to scale. - mask_obs - Restrict both the derivation of scaling parameters and the scaling itself - to a certain set of observations. The mask is specified as a boolean array - or a string referring to an array in :attr:`~anndata.AnnData.obs`. - This will transform data from csc to csr format if `issparse(data)`. + {mask} Returns ------- @@ -142,13 +155,19 @@ def scale[A: _Array]( msg = f"`obsm` argument inappropriate for value of type {type(data)}" raise ValueError(msg) return scale_array( - data, zero_center=zero_center, max_value=max_value, copy=copy, mask_obs=mask_obs + data, + zero_center=zero_center, + max_value=max_value, + copy=copy, + mask=mask, + mask_obs=mask_obs, ) @scale.register(np.ndarray) @scale.register(DaskArray) @scale.register(CSBase) +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) def scale_array[A: _Array]( x: A, *, @@ -156,6 +175,7 @@ def scale_array[A: _Array]( max_value: float | None = None, copy: bool = False, return_mean_std: bool = False, + mask: NDArray[np.bool] | None = None, mask_obs: NDArray[np.bool] | None = None, ) -> ( A @@ -185,17 +205,18 @@ def scale_array[A: _Array]( ) x = x.astype(np.float64) - mask_obs = ( + mask = _mask_arg(mask, mask_obs, dim="obs") + mask = ( # For CSR matrices, default to a set mask to take the `scale_array_masked` path. # This is faster than the maskless `axis_mul_or_truediv` path. np.ones(x.shape[0], dtype=np.bool) - if isinstance(x, CSRBase) and mask_obs is None and not zero_center - else _check_mask(x, mask_obs, "obs") + if isinstance(x, CSRBase) and mask is None and not zero_center + else _check_mask(x, mask, "obs") ) - if mask_obs is not None: + if mask is not None: return scale_array_masked( x, - mask_obs, + mask, zero_center=zero_center, max_value=max_value, return_mean_std=return_mean_std, @@ -294,6 +315,7 @@ def scale_and_clip_csr( @scale.register(AnnData) +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) def scale_anndata( adata: AnnData, *, @@ -302,16 +324,18 @@ def scale_anndata( copy: bool = False, layer: str | None = None, obsm: str | None = None, - mask_obs: NDArray[np.bool] | str | None = None, + mask: Mask | None = None, + mask_obs: Mask | None = None, ) -> AnnData | None: adata = adata.copy() if copy else adata + mask = _mask_arg(mask, mask_obs, dim="obs") str_mean_std = ("mean", "std") - if mask_obs is not None: - if isinstance(mask_obs, str): - str_mean_std = (f"mean of {mask_obs}", f"std of {mask_obs}") + if mask is not None: + if isinstance(mask, str | AdRef): + str_mean_std = (f"mean of {mask}", f"std of {mask}") else: str_mean_std = ("mean with mask", "std with mask") - mask_obs = _check_mask(adata, mask_obs, "obs") + mask = _check_mask(adata, mask, "obs") view_to_actual(adata) x = _get_arr(adata, layer=layer, obsm=obsm) raise_not_implemented_error_if_backed_type(x, "scale") @@ -321,7 +345,7 @@ def scale_anndata( max_value=max_value, copy=False, # because a copy has already been made, if it were to be made return_mean_std=True, - mask_obs=mask_obs, + mask=mask, ) _set_obs_rep(adata, x, layer=layer, obsm=obsm) return adata if copy else None diff --git a/src/scanpy/preprocessing/_simple.py b/src/scanpy/preprocessing/_simple.py index 935a2e79c6..3d014d60a0 100644 --- a/src/scanpy/preprocessing/_simple.py +++ b/src/scanpy/preprocessing/_simple.py @@ -24,7 +24,7 @@ from .. import logging as logg from .._compat import CSBase, CSRBase, DaskArray -from .._docs import doc_rng +from .._docs import doc_ref_compat, doc_rng from .._settings import settings from .._utils import ( _doc_params, @@ -48,6 +48,7 @@ from numpy.typing import NDArray from .._utils.random import RNGLike, SeedLike + from ..get.get import Mask def filter_cells( @@ -691,7 +692,7 @@ def sample( copy: Literal[False] = False, replace: bool = False, axis: Literal["obs", 0, "var", 1] = "obs", - p: str | NDArray[np.bool] | NDArray[np.floating] | None = None, + p: Mask | NDArray[np.floating] | None = None, ) -> None: ... @overload def sample( @@ -703,7 +704,7 @@ def sample( copy: Literal[True], replace: bool = False, axis: Literal["obs", 0, "var", 1] = "obs", - p: str | NDArray[np.bool] | NDArray[np.floating] | None = None, + p: Mask | NDArray[np.floating] | None = None, ) -> AnnData: ... @overload def sample[A: np.ndarray | CSBase | DaskArray]( @@ -715,11 +716,11 @@ def sample[A: np.ndarray | CSBase | DaskArray]( copy: bool = False, replace: bool = False, axis: Literal["obs", 0, "var", 1] = "obs", - p: str | NDArray[np.bool] | NDArray[np.floating] | None = None, + p: Mask | NDArray[np.floating] | None = None, ) -> tuple[A, NDArray[np.int64]]: ... -@_doc_params(rng=doc_rng) +@_doc_params(rng=doc_rng, ref=doc_ref_compat) def sample( # noqa: PLR0912 data: AnnData | np.ndarray | CSBase | DaskArray, fraction: float | None = None, @@ -729,7 +730,7 @@ def sample( # noqa: PLR0912 copy: bool = False, replace: bool = False, axis: Literal["obs", 0, "var", 1] = "obs", - p: str | NDArray[np.bool] | NDArray[np.floating] | None = None, + p: Mask | NDArray[np.floating] | None = None, ) -> AnnData | tuple[np.ndarray | CSBase | DaskArray, NDArray[np.int64]] | None: r"""Sample observations or variables with or without replacement. @@ -757,7 +758,8 @@ def sample( # noqa: PLR0912 Sample `obs`\ ervations (axis 0) or `var`\ iables (axis 1). p Drawing probabilities (floats) or mask (bools). - Either an `axis`-sized array, or the name of a column. + Either an `axis`-sized array, or a reference to one, e.g. `A.obs['is_control']`. + {ref} If `p` is an array of probabilities, it must sum to 1. Returns diff --git a/src/scanpy/tools/_ingest.py b/src/scanpy/tools/_ingest.py index 31f8ed5061..d0ace4790c 100644 --- a/src/scanpy/tools/_ingest.py +++ b/src/scanpy/tools/_ingest.py @@ -22,7 +22,7 @@ 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 ..get.get import MultiAcc, _ref_from_json, _rep_from_json from ..neighbors import FlatTree from ._utils import _choose_representation_compat @@ -350,7 +350,7 @@ def _init_neighbors(self, adata: AnnData, neighbors_key: str | None) -> None: def _init_pca(self, adata: AnnData) -> None: self._pca_centered = adata.uns["pca"]["params"]["zero_center"] self._pca_mask = _check_mask( - adata, adata.uns["pca"]["params"]["mask_var"], "var" + adata, _ref_from_json(adata.uns["pca"]["params"]["mask_var"]), "var" ) if self._pca_mask is not None: diff --git a/src/scanpy/tools/_rank_genes_groups.py b/src/scanpy/tools/_rank_genes_groups.py index 4f45761e69..a82693ba7b 100644 --- a/src/scanpy/tools/_rank_genes_groups.py +++ b/src/scanpy/tools/_rank_genes_groups.py @@ -10,21 +10,25 @@ from anndata import AnnData from fast_array_utils.numba import njit from scipy import sparse +from scverse_misc import Deprecation, deprecated_arg from .. import _utils from .. import logging as logg from .._compat import CSBase, DaskArray, warn +from .._docs import doc_mask from .._settings import Default, Preset, settings from .._settings.presets import DETest from .._utils import ( + _doc_params, _numba_thread_limit, check_nonnegative_integers, + dim_acc, get_literal_vals, - obs_acc, raise_not_implemented_error_if_backed_type, ) 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 @@ -32,6 +36,8 @@ from numpy.typing import NDArray + from ..get.get import Mask + type _CorrMethod = Literal["benjamini-hochberg", "bonferroni"] type _TestResult = tuple[int, NDArray[np.floating], NDArray[np.floating] | None] @@ -365,7 +371,7 @@ def _aggregate_group_stats( index=pd.RangeIndex(len(codes)).astype(str), ), ) - out = aggregate(agg_adata, by=obs_acc("_g"), func=funcs, dof=1) + out = aggregate(agg_adata, by=dim_acc("_g"), func=funcs, dof=1) idx = out.obs_names.astype(int).to_numpy() mean[idx] = np.asarray(out.layers["mean"]) if need_var: @@ -740,13 +746,15 @@ def _build_stats_dataframe( return df +@_doc_params( + mask=doc_mask("Select subset of genes to use in statistical tests.", dim="var") +) +@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead.")) def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 adata: AnnData, groupby: str, *, - mask_var: NDArray[np.bool] | str | Default | None = Default( - preset=("rank_genes_groups", "mask_var") - ), + mask: Mask | Default | None = Default(preset=("rank_genes_groups", "mask")), use_raw: bool | None = None, groups: Literal["all"] | Iterable[str] = "all", reference: str = "rest", @@ -762,6 +770,7 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 mean_in_log_space: bool | Default = Default( preset=("rank_genes_groups", "mean_in_log_space") ), + mask_var: Mask | None = None, **kwds, ) -> AnnData | None: r"""Rank genes for characterizing groups. @@ -785,8 +794,7 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 Annotated data matrix. groupby The key of the observations grouping to consider. - mask_var - Select subset of genes to use in statistical tests. + {mask} use_raw Use `raw` attribute of `adata` if present. The default behavior is to use `raw` if present. layer @@ -826,8 +834,8 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 copy Whether to copy `adata` or modify it inplace. mean_in_log_space - Whether to do :math:`\log(\operatorname{mean}(e^x))` (`False`) - or :math:`\log(e^{\operatorname{mean}(x)})` (`True`). + Whether to do :math:`\log(\operatorname{{mean}}(e^x))` (`False`) + or :math:`\log(e^{{\operatorname{{mean}}(x)}})` (`True`). The former is accurate, while the latter is a faster approximation that underestimates this accurate result in the presence of many outliers. kwds @@ -878,8 +886,9 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 >>> sc.pl.rank_genes_groups(adata) """ - if isinstance(mask_var, Default): - mask_var = settings.preset.rank_genes_groups.mask_var + mask = _mask_arg(mask, mask_var, dim="var") + if isinstance(mask, Default): + mask = settings.preset.rank_genes_groups.mask 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. @@ -895,7 +904,7 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 ) warn(msg, DeprecationWarning) - mask_var = _check_mask(adata, mask_var, "var") + mask_var = _check_mask(adata, mask, "var") if use_raw is None: use_raw = adata.raw is not None diff --git a/tests/test_ingest.py b/tests/test_ingest.py index 70f7a564c5..3c7df19874 100644 --- a/tests/test_ingest.py +++ b/tests/test_ingest.py @@ -90,6 +90,7 @@ def test_representation_acc(adatas) -> None: np.testing.assert_array_equal(ing._obsm["rep"], adata_new.obsm["X_pca"]) +@needs.anndata_acc @pytest.mark.parametrize("as_sparse", [False, True]) def test_pca_transform_uses_reference_mean( as_sparse, monkeypatch: pytest.MonkeyPatch @@ -108,7 +109,7 @@ def test_pca_transform_uses_reference_mean( adata_ref = sc.AnnData(ref_x) adata_new = sc.AnnData(query_x) adata_ref.var["selected"] = mask - sc.pp.pca(adata_ref, n_comps=3, mask_var="selected") + sc.pp.pca(adata_ref, n_comps=3, mask="var.selected") sc.pp.neighbors(adata_ref, n_neighbors=3, n_pcs=2) ing = sc.tl.Ingest(adata_ref) diff --git a/tests/test_normalization.py b/tests/test_normalization.py index 1058b737e9..0cb544a743 100644 --- a/tests/test_normalization.py +++ b/tests/test_normalization.py @@ -17,6 +17,7 @@ check_rep_mutation, check_rep_results, ) +from testing.scanpy._pytest.marks import needs # TODO: Add support for sparse-in-dask from testing.scanpy._pytest.params import ARRAY_TYPES, ARRAY_TYPES_DENSE @@ -215,8 +216,14 @@ def _check_pearson_pca_fields(ad, n_cells, n_comps): [ pytest.param(False, dict(), "n_genes", id="no_hvg"), pytest.param(True, dict(), "n_hvgs", id="hvg_default"), - pytest.param(True, dict(mask_var=None), "n_genes", id="hvg_opt_out"), - pytest.param(False, dict(mask_var="test_mask"), "n_unmasked", id="mask"), + pytest.param(True, dict(mask=None), "n_genes", id="hvg_opt_out"), + pytest.param( + False, + dict(mask="var.test_mask"), + "n_unmasked", + id="mask", + marks=needs.anndata_acc, + ), ], ) def test_normalize_pearson_residuals_pca( diff --git a/tests/test_pca.py b/tests/test_pca.py index 2a063b0af5..a1fedface0 100644 --- a/tests/test_pca.py +++ b/tests/test_pca.py @@ -402,15 +402,16 @@ def test_pca_n_pcs(): # We use all possible array types here since this error should be raised before # PCA can realize that it got a Dask array +@needs.anndata_acc @pytest.mark.parametrize("array_type", ARRAY_TYPES_ALL) def test_mask_var_error(array_type): - """Check if mask_var="..." throws an error if the annotation is missing.""" + """Check if mask="..." throws an error if the annotation is missing.""" adata = AnnData(array_type(A_list).astype("float32")) with pytest.raises( ValueError, - match=r"Did not find `adata\.var\['highly_variable'\]`\.", + match=r"Did not find `A\.var\['highly_variable'\]` in `adata`\.", ): - sc.pp.pca(adata, mask_var="highly_variable") + sc.pp.pca(adata, mask="var.highly_variable") def test_mask_length_error(): @@ -420,22 +421,26 @@ def test_mask_length_error(): with pytest.raises( ValueError, match=r"The shape of the mask do not match the data\." ): - sc.pp.pca(adata, mask_var=mask_var, copy=True) + sc.pp.pca(adata, mask=mask_var, copy=True) -@pytest.mark.parametrize("mask_type", ["highly_variable", "array"]) -def test_obsm_mask_error(mask_type: Literal["highly_variable", "array"]) -> None: - """Check that trying to use mask_var with obsm raises an error.""" +@pytest.mark.parametrize( + "mask_type", + [pytest.param("var.highly_variable", marks=needs.anndata_acc), "array"], +) +def test_obsm_mask_error(mask_type: Literal["var.highly_variable", "array"]) -> None: + """Check that trying to use mask with obsm raises an error.""" adata = AnnData(A_list) mask_var = ( _helpers.random_mask(adata.shape[1]) if mask_type == "array" else mask_type ) with pytest.raises( - ValueError, match=r"Argument `mask_var` is incompatible with `obsm`." + ValueError, match=r"Argument `mask` is incompatible with `obsm`." ): - sc.pp.pca(adata, mask_var=mask_var, obsm="X_pca", copy=True) + sc.pp.pca(adata, mask=mask_var, obsm="X_pca", copy=True) +@needs.anndata_acc def test_mask_var_argument_equivalence(float_dtype, array_type): """Test if pca result is equal when given mask as boolarray vs string.""" rng = np.random.default_rng() @@ -443,11 +448,11 @@ def test_mask_var_argument_equivalence(float_dtype, array_type): mask_var = _helpers.random_mask(adata_base.shape[1], rng=rng) adata = adata_base.copy() - sc.pp.pca(adata, mask_var=mask_var, dtype=float_dtype) + sc.pp.pca(adata, mask=mask_var, dtype=float_dtype) adata_w_mask = adata_base.copy() adata_w_mask.var["mask"] = mask_var - sc.pp.pca(adata_w_mask, mask_var="mask", dtype=float_dtype) + sc.pp.pca(adata_w_mask, mask="var.mask", dtype=float_dtype) adata, adata_w_mask = map(AnnData.to_memory, [adata, adata_w_mask]) assert np.allclose( @@ -468,7 +473,7 @@ def test_mask(request: pytest.FixtureRequest, array_type): mask_var = _helpers.random_mask(adata.shape[1]) adata_masked = adata[:, mask_var].copy() - sc.pp.pca(adata, mask_var=mask_var) + sc.pp.pca(adata, mask=mask_var) sc.pp.pca(adata_masked) masked_var_loadings = adata.varm["PCs"][~mask_var] @@ -501,7 +506,7 @@ def test_mask_defaults(array_type, float_dtype): without_var, with_var = map(AnnData.to_memory, [without_var, with_var]) assert not np.array_equal(without_var.obsm["X_pca"], with_var.obsm["X_pca"]) - with_no_mask = sc.pp.pca(adata, mask_var=None, copy=True, dtype=float_dtype) + with_no_mask = sc.pp.pca(adata, mask=None, copy=True, dtype=float_dtype) with_no_mask = with_no_mask.to_memory() assert np.array_equal(without_var.obsm["X_pca"], with_no_mask.obsm["X_pca"]) @@ -523,8 +528,8 @@ def test_pca_rep(rep: Literal["layer", "obsm"]) -> None: pytest.fail(f"Unknown {rep=}") del rep_adata.X - sc.pp.pca(adata, mask_var=None) - sc.pp.pca(rep_adata, **{rep: "counts"}, mask_var=None) + sc.pp.pca(adata, mask=None) + sc.pp.pca(rep_adata, **{rep: "counts"}, mask=None) assert rep_adata.uns["pca"]["params"][rep] == "counts" assert rep not in adata.uns["pca"]["params"] diff --git a/tests/test_rank_genes_groups.py b/tests/test_rank_genes_groups.py index 6caabad6ca..dc430f19e3 100644 --- a/tests/test_rank_genes_groups.py +++ b/tests/test_rank_genes_groups.py @@ -346,7 +346,7 @@ def test_mask_n_genes(n_genes_add, n_genes_out_add): rank_genes_groups( pbmc, - mask_var=mask_var, + mask=mask_var, groupby="bulk_labels", groups=["CD14+ Monocyte", "Dendritic"], reference="CD14+ Monocyte", @@ -375,7 +375,7 @@ def test_mask_not_equal(): run(n_genes=n_genes) no_mask = pbmc.uns["rank_genes_groups"]["names"] - run(mask_var=mask_var) + run(mask=mask_var) with_mask = pbmc.uns["rank_genes_groups"]["names"] assert not np.array_equal(no_mask, with_mask) diff --git a/tests/test_scaling.py b/tests/test_scaling.py index 61ee930067..8e2064b4c4 100644 --- a/tests/test_scaling.py +++ b/tests/test_scaling.py @@ -9,6 +9,7 @@ from scipy import sparse import scanpy as sc +from testing.scanpy._pytest.marks import needs # test "data" for 3 cells * 4 genes X_original = [ @@ -80,7 +81,7 @@ @pytest.mark.parametrize("dtype", [np.float32, np.int64]) @pytest.mark.parametrize("zero_center", [True, False], ids=["center", "no_center"]) @pytest.mark.parametrize( - ("mask_obs", "x", "x_centered", "x_scaled"), + ("mask", "x", "x_centered", "x_scaled"), [ pytest.param( None, X_original, X_centered_original, X_scaled_original, id="no_mask" @@ -94,9 +95,7 @@ ), ], ) -def test_scale( - *, typ, container, zero_center, dtype, mask_obs, x, x_centered, x_scaled -): +def test_scale(*, typ, container, zero_center, dtype, mask, x, x_centered, x_scaled): x = AnnData(typ(x, dtype=dtype)) if container == "anndata" else typ(x, dtype=dtype) with warnings.catch_warnings(): # TODO: fix setting slices of sparse matrices in scale() @@ -108,7 +107,7 @@ def test_scale( else nullcontext() ): scaled = sc.pp.scale( - x, zero_center=zero_center, copy=container == "array", mask_obs=mask_obs + x, zero_center=zero_center, copy=container == "array", mask=mask ) received = sparse.csr_matrix( # noqa: TID251 x.X if scaled is None else scaled @@ -117,14 +116,64 @@ def test_scale( assert np.allclose(received, expected) -def test_mask_string(): - with pytest.raises(ValueError, match=r"Cannot.*refer.*mask.*without.*anndata"): - sc.pp.scale(np.array(X_original), mask_obs="mask") +@pytest.fixture +def adata_masked() -> AnnData: adata = AnnData(np.array(X_for_mask, dtype="float32")) adata.obs["some cells"] = np.array((0, 0, 1, 1, 1, 0, 0), dtype=bool) - sc.pp.scale(adata, mask_obs="some cells") - assert np.array_equal(adata.X, X_centered_for_mask) - assert "mean of some cells" in adata.var.columns + return adata + + +@needs.anndata_acc +@pytest.mark.parametrize("as_ref", [True, False], ids=["ref", "str"]) +def test_mask_ref(adata_masked: AnnData, *, as_ref: bool) -> None: + from anndata.acc import A + + mask = A.obs["some cells"] if as_ref else "obs.some cells" + with pytest.raises(ValueError, match=r"Cannot.*refer.*mask.*without.*anndata"): + sc.pp.scale(np.array(X_original), mask=mask) + sc.pp.scale(adata_masked, mask=mask) + assert np.array_equal(adata_masked.X, X_centered_for_mask) + assert "mean of A.obs['some cells']" in adata_masked.var.columns + + +@needs.anndata_acc +def test_mask_wrong_dim(adata_masked: AnnData) -> None: + with pytest.raises(ValueError, match=r"Dimension of .* \(var\) does not match"): + sc.pp.scale(adata_masked, mask="var.some genes") + + +@pytest.mark.parametrize( + ("preset", "col"), + [ + pytest.param(sc.Preset.ScanpyV1, "mean of some cells", id="v1"), + pytest.param( + sc.Preset.ScanpyV2Preview, + "mean of A.obs['some cells']", + id="v2", + marks=needs.anndata_acc, + ), + ], +) +def test_mask_obs_deprecated( + adata_masked: AnnData, *, preset: sc.Preset, col: str +) -> None: + r"""The deprecated `mask_obs` interprets :class:`str`\ s as `.obs` column names.""" + with ( + sc.settings.override(preset=preset), + pytest.warns(FutureWarning, match=r"argument mask_obs is deprecated"), + ): + sc.pp.scale(adata_masked, zero_center=True, mask_obs="some cells") + assert np.array_equal(adata_masked.X, X_centered_for_mask) + assert col in adata_masked.var.columns + + +def test_mask_both() -> None: + mask = np.array((0, 0, 1, 1, 1, 0, 0), dtype=bool) + with ( + pytest.warns(FutureWarning, match=r"argument mask_obs is deprecated"), + pytest.raises(TypeError, match=r"Pass either `mask` or `mask_obs`, not both"), + ): + sc.pp.scale(np.array(X_for_mask, dtype="float32"), mask=mask, mask_obs=mask) @pytest.mark.parametrize("zero_center", [True, False], ids=["center", "no_center"]) @@ -142,7 +191,7 @@ def test_clip(*, zero_center: bool) -> None: @pytest.mark.parametrize( - ("mask_obs", "x", "x_scaled", "x_clipped"), + ("mask", "x", "x_scaled", "x_clipped"), [ pytest.param( None, X_original, X_scaled_original, X_scaled_original_clipped, id="no_mask" @@ -157,8 +206,8 @@ def test_clip(*, zero_center: bool) -> None: ], ) @pytest.mark.parametrize("clip", [False, True], ids=["no_clip", "clip"]) -def test_scale_sparse(*, mask_obs, x, x_scaled, x_clipped, clip): +def test_scale_sparse(*, mask, x, x_scaled, x_clipped, clip): max_value, expected = (1, x_clipped) if clip else (None, x_scaled) adata = AnnData(sparse.csr_matrix(x).astype(np.float32)) # noqa: TID251 - sc.pp.scale(adata, mask_obs=mask_obs, zero_center=False, max_value=max_value) + sc.pp.scale(adata, mask=mask, zero_center=False, max_value=max_value) assert np.allclose(sparse.csr_matrix(adata.X).toarray(), expected) # noqa: TID251 From 3a5cd552e702aa0e9f3ec8156582debf4756ecee Mon Sep 17 00:00:00 2001 From: "Philipp A." Date: Tue, 1 Sep 2026 17:49:51 +0200 Subject: [PATCH 2/7] relnote --- docs/release-notes/{4331.feat.md => 4333.feat.md} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename docs/release-notes/{4331.feat.md => 4333.feat.md} (100%) diff --git a/docs/release-notes/4331.feat.md b/docs/release-notes/4333.feat.md similarity index 100% rename from docs/release-notes/4331.feat.md rename to docs/release-notes/4333.feat.md From 2c7ab0224deeed5f64dfb0e1fd85e94304e393ad Mon Sep 17 00:00:00 2001 From: "Philipp A." Date: Tue, 1 Sep 2026 22:05:48 +0200 Subject: [PATCH 3/7] revert hack --- pyproject.toml | 2 +- src/scanpy/get/get.py | 6 ++---- 2 files changed, 3 insertions(+), 5 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index ab57866b81..422765fc13 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -127,7 +127,7 @@ docs = [ "scanpydoc>=0.17.4", "scverse-misc[sphinx]", "sphinx>=9.1", - "sphinx-autodoc-typehints>=1.25.2", + "sphinx-autodoc-typehints>=3.13.5", "sphinx-book-theme>=1.1", "sphinx-copybutton", "sphinx-design", diff --git a/src/scanpy/get/get.py b/src/scanpy/get/get.py index 4582ab94ef..2653605231 100644 --- a/src/scanpy/get/get.py +++ b/src/scanpy/get/get.py @@ -21,7 +21,7 @@ from collections.abc import Iterable from typing import Any, Literal, Unpack - from anndata.acc import RefAcc + from anndata.acc import Idx2D, RefAcc from .._compat import DaskArray @@ -32,12 +32,10 @@ if TYPE_CHECKING or find_spec("anndata.acc"): - from anndata.acc import AdRef, GraphAcc, Idx2D, LayerAcc, MultiAcc + from anndata.acc import AdRef, GraphAcc, LayerAcc, MultiAcc else: AdRef = type("AdRef", (), dict(__module__="anndata.acc")) GraphAcc = type("GraphAcc", (), dict(__module__="anndata.acc")) - # https://github.com/tox-dev/sphinx-autodoc-typehints/issues/764 - type Idx2D = object LayerAcc = type("LayerAcc", (), dict(__module__="anndata.acc")) MultiAcc = type("MultiAcc", (), dict(__module__="anndata.acc")) From 8d909d68f052ced7938c1dba0d0edac6899d4305 Mon Sep 17 00:00:00 2001 From: "Philipp A." Date: Tue, 1 Sep 2026 22:15:15 +0200 Subject: [PATCH 4/7] whoops --- tests/test_scaling.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_scaling.py b/tests/test_scaling.py index 8e2064b4c4..c75a075eeb 100644 --- a/tests/test_scaling.py +++ b/tests/test_scaling.py @@ -150,7 +150,7 @@ def test_mask_wrong_dim(adata_masked: AnnData) -> None: sc.Preset.ScanpyV2Preview, "mean of A.obs['some cells']", id="v2", - marks=needs.anndata_acc, + marks=needs.scanpy2, ), ], ) From dd09b9f4c550d0f1e61b7028338ff85ffefb1aef Mon Sep 17 00:00:00 2001 From: Phil Schaf Date: Thu, 3 Sep 2026 16:43:22 +0200 Subject: [PATCH 5/7] simplify a little --- src/scanpy/experimental/pp/_normalization.py | 15 +++++---------- src/scanpy/get/get.py | 15 +++++++++++++++ src/scanpy/preprocessing/_pca/__init__.py | 16 ++++++---------- src/scanpy/preprocessing/_scale.py | 9 ++++++--- tests/test_scaling.py | 12 ++++++++++++ 5 files changed, 44 insertions(+), 23 deletions(-) diff --git a/src/scanpy/experimental/pp/_normalization.py b/src/scanpy/experimental/pp/_normalization.py index a0f34e7f4d..b429582305 100644 --- a/src/scanpy/experimental/pp/_normalization.py +++ b/src/scanpy/experimental/pp/_normalization.py @@ -15,7 +15,7 @@ from ..._compat import CSBase, warn from ..._keys import _embedding_keys from ..._settings import Default -from ..._utils import _doc_params, check_nonnegative_integers, dim_acc, view_to_actual +from ..._utils import _doc_params, check_nonnegative_integers, view_to_actual from ..._utils.random import _accepts_legacy_random_state from ...experimental._docs import ( doc_adata, @@ -27,7 +27,7 @@ doc_pca_chunk, ) from ...get import _check_mask, _get_arr, _set_obs_rep -from ...get.get import _mask_arg +from ...get.get import _mask_arg, _mask_hvg from ...preprocessing._docs import doc_mask_var from ...preprocessing._pca import pca @@ -171,7 +171,8 @@ def normalize_pearson_residuals( inplace=doc_inplace, ) @_accepts_legacy_random_state(0) -@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead.")) +# `stacklevel=2` skips `_accepts_legacy_random_state`’s wrapper frame +@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def normalize_pearson_residuals_pca( adata: AnnData, *, @@ -231,13 +232,7 @@ def normalize_pearson_residuals_pca( """ key_added = kwargs_pca.get("key_added", settings.preset.pca.key_added) keys = _embedding_keys("pca", key_added) - mask = _mask_arg(mask, mask_var, dim="var") - if isinstance(mask, Default): - mask = ( - dim_acc("highly_variable", dim="var") - if "highly_variable" in adata.var - else None - ) + mask = _mask_hvg(adata, _mask_arg(mask, mask_var, dim="var")) mask_var = _check_mask(adata, mask, "var") if mask_var is not None: diff --git a/src/scanpy/get/get.py b/src/scanpy/get/get.py index 2653605231..b33d44634b 100644 --- a/src/scanpy/get/get.py +++ b/src/scanpy/get/get.py @@ -621,12 +621,27 @@ def _mask_arg[M]( raise TypeError(msg) return dim_acc(legacy, dim=dim) if isinstance(legacy, str) else legacy if isinstance(mask, str): + if not find_spec("anndata.acc"): + msg = ( + f"`mask={mask!r}` requires `anndata>=0.13.3`. " + f"Pass a boolean array instead, or a column name via the deprecated `mask_{dim}`." + ) + raise ImportError(msg) from anndata.acc import A return A.resolve(mask, vec=True) return mask +def _mask_hvg[M](adata: AnnData, mask: M | Default) -> M | None: + """Resolve a `Default` `mask` to `.var['highly_variable']` if there is one.""" + if not isinstance(mask, Default): + return mask + if "highly_variable" not in adata.var: + return None + return dim_acc("highly_variable", dim="var") + + def _check_mask[M: NDArray[np.bool] | NDArray[np.floating] | pd.Series | None]( data: AnnData | np.ndarray | CSBase | DaskArray, mask: str | AdRef[Idx2D | int, AnnData] | M, diff --git a/src/scanpy/preprocessing/_pca/__init__.py b/src/scanpy/preprocessing/_pca/__init__.py index 39ebd954a5..7a8e336078 100644 --- a/src/scanpy/preprocessing/_pca/__init__.py +++ b/src/scanpy/preprocessing/_pca/__init__.py @@ -11,10 +11,10 @@ from ..._docs import doc_rng from ..._keys import _embedding_keys from ..._settings import Default, settings -from ..._utils import _doc_params, dim_acc, get_literal_vals, is_backed_type +from ..._utils import _doc_params, get_literal_vals, is_backed_type from ..._utils.random import _accepts_legacy_random_state, _legacy_random_state from ...get import _check_mask, _get_arr -from ...get.get import _mask_arg, _ref_to_json +from ...get.get import _mask_arg, _mask_hvg, _ref_to_json from .._docs import doc_mask_var from ._compat import _pca_compat_sparse @@ -53,7 +53,8 @@ @_doc_params(mask=doc_mask_var, rng=doc_rng) @_accepts_legacy_random_state(0) -@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead.")) +# `stacklevel=2` skips `_accepts_legacy_random_state`’s wrapper frame +@deprecated_arg("mask_var", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def pca( # noqa: PLR0912, PLR0913, PLR0915 data: AnnData | np.ndarray | CSBase, n_comps: int | None = None, @@ -222,15 +223,10 @@ def pca( # noqa: PLR0912, PLR0913, PLR0915 adata = AnnData(data) mask = _mask_arg(mask, mask_var, dim="var") - if isinstance(mask, Default): - mask = ( - dim_acc("highly_variable", dim="var") - if "highly_variable" in adata.var - else None - ) - elif mask is not None and obsm is not None: + if not isinstance(mask, Default) and mask is not None and obsm is not None: msg = "Argument `mask` is incompatible with `obsm`." raise ValueError(msg) + mask = _mask_hvg(adata, mask) mask_param, mask_var = mask, _check_mask(adata, mask, "var") adata_comp = adata[:, mask_var] if mask_var is not None else adata diff --git a/src/scanpy/preprocessing/_scale.py b/src/scanpy/preprocessing/_scale.py index 7647e075c6..a94b18cc3b 100644 --- a/src/scanpy/preprocessing/_scale.py +++ b/src/scanpy/preprocessing/_scale.py @@ -83,7 +83,8 @@ def clip_array( ) ) @singledispatch -@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) +# `stacklevel=2` skips `singledispatch`’s dispatcher frame +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def scale[A: _Array]( data: AnnData | A, *, @@ -167,7 +168,8 @@ def scale[A: _Array]( @scale.register(np.ndarray) @scale.register(DaskArray) @scale.register(CSBase) -@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) +# `stacklevel=2` skips `singledispatch`’s dispatcher frame +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def scale_array[A: _Array]( x: A, *, @@ -315,7 +317,8 @@ def scale_and_clip_csr( @scale.register(AnnData) -@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead.")) +# `stacklevel=2` skips `singledispatch`’s dispatcher frame +@deprecated_arg("mask_obs", Deprecation("1.13.0", "Use `mask` instead."), stacklevel=2) def scale_anndata( adata: AnnData, *, diff --git a/tests/test_scaling.py b/tests/test_scaling.py index c75a075eeb..f114807c0a 100644 --- a/tests/test_scaling.py +++ b/tests/test_scaling.py @@ -167,6 +167,18 @@ def test_mask_obs_deprecated( assert col in adata_masked.var.columns +def test_mask_obs_deprecated_fallback() -> None: + # extra test for the singledispatch’s fallback branch calling `scale_array` + # (registered types like `np.ndarray` never run it) + mask = np.array((0, 0, 1, 1, 1, 0, 0), dtype=bool) + scale_fallback = sc.pp.scale.dispatch(object) + with pytest.warns(FutureWarning, match=r"argument mask_obs is deprecated"): + scaled = scale_fallback( + np.array(X_for_mask, dtype="float32"), copy=True, mask_obs=mask + ) + assert np.array_equal(scaled, X_centered_for_mask) + + def test_mask_both() -> None: mask = np.array((0, 0, 1, 1, 1, 0, 0), dtype=bool) with ( From 7bd0a9451e33bcfe68c4bb176f3ecf4a201d0a09 Mon Sep 17 00:00:00 2001 From: Phil Schaf Date: Mon, 7 Sep 2026 13:46:14 +0200 Subject: [PATCH 6/7] simplify rgg mask --- src/scanpy/_settings/presets.py | 5 ++--- src/scanpy/tools/_rank_genes_groups.py | 4 +--- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/src/scanpy/_settings/presets.py b/src/scanpy/_settings/presets.py index 1e3fe62cb1..abf6fd4dd2 100644 --- a/src/scanpy/_settings/presets.py +++ b/src/scanpy/_settings/presets.py @@ -93,7 +93,6 @@ class BasicEmbeddingPreset(NamedTuple): class RankGenesGroupsPreset(NamedTuple): method: DETest - mask: str | None mean_in_log_space: bool @@ -245,10 +244,10 @@ def rank_genes_groups() -> Mapping[Preset, RankGenesGroupsPreset]: """ return { Preset.ScanpyV1: RankGenesGroupsPreset( - method="t-test", mask=None, mean_in_log_space=True + method="t-test", mean_in_log_space=True ), Preset.ScanpyV2Preview: RankGenesGroupsPreset( - method="wilcoxon", mask=None, mean_in_log_space=False + method="wilcoxon", mean_in_log_space=False ), } diff --git a/src/scanpy/tools/_rank_genes_groups.py b/src/scanpy/tools/_rank_genes_groups.py index a82693ba7b..c7db26099a 100644 --- a/src/scanpy/tools/_rank_genes_groups.py +++ b/src/scanpy/tools/_rank_genes_groups.py @@ -754,7 +754,7 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 adata: AnnData, groupby: str, *, - mask: Mask | Default | None = Default(preset=("rank_genes_groups", "mask")), + mask: Mask | None = None, use_raw: bool | None = None, groups: Literal["all"] | Iterable[str] = "all", reference: str = "rest", @@ -887,8 +887,6 @@ def rank_genes_groups( # noqa: PLR0912, PLR0913, PLR0915 """ mask = _mask_arg(mask, mask_var, dim="var") - if isinstance(mask, Default): - mask = settings.preset.rank_genes_groups.mask 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. From a2729d89b5e7f61a1d9fcb41fb515a3a16742cd7 Mon Sep 17 00:00:00 2001 From: Phil Schaf Date: Mon, 7 Sep 2026 13:48:38 +0200 Subject: [PATCH 7/7] remove dim_acc default --- src/scanpy/_utils/__init__.py | 2 +- src/scanpy/preprocessing/_highly_variable_genes.py | 4 +++- src/scanpy/tools/_rank_genes_groups.py | 2 +- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/src/scanpy/_utils/__init__.py b/src/scanpy/_utils/__init__.py index dc7a9a8faf..ddc04fc2af 100644 --- a/src/scanpy/_utils/__init__.py +++ b/src/scanpy/_utils/__init__.py @@ -988,7 +988,7 @@ def _resolve_axis( raise ValueError(msg) -def dim_acc(col: str, *, dim: Literal["obs", "var"] = "obs") -> str | AdRef: +def dim_acc(col: str, *, dim: Literal["obs", "var"]) -> str | AdRef: """Get reference to the `col`umn of `adata.{dim}` the way the active preset expects it.""" from .._settings import Preset, settings diff --git a/src/scanpy/preprocessing/_highly_variable_genes.py b/src/scanpy/preprocessing/_highly_variable_genes.py index d76e76cfd3..eb9207e7fe 100644 --- a/src/scanpy/preprocessing/_highly_variable_genes.py +++ b/src/scanpy/preprocessing/_highly_variable_genes.py @@ -181,7 +181,9 @@ def _highly_variable_genes_seurat_v3( # noqa: PLR0912, PLR0915 ) if batch_key is not None: aggregated_mean_var = aggregate( - adata_agg, by=dim_acc("__hvg_v3_batch_info__"), func=["mean", "var"] + adata_agg, + by=dim_acc("__hvg_v3_batch_info__", dim="obs"), + func=["mean", "var"], ) aggregated_mean_var.layers["mean"], aggregated_mean_var.layers["var"] = ( materialize_as_ndarray( diff --git a/src/scanpy/tools/_rank_genes_groups.py b/src/scanpy/tools/_rank_genes_groups.py index c7db26099a..4c4ed63614 100644 --- a/src/scanpy/tools/_rank_genes_groups.py +++ b/src/scanpy/tools/_rank_genes_groups.py @@ -371,7 +371,7 @@ def _aggregate_group_stats( index=pd.RangeIndex(len(codes)).astype(str), ), ) - out = aggregate(agg_adata, by=dim_acc("_g"), func=funcs, dof=1) + out = aggregate(agg_adata, by=dim_acc("_g", dim="obs"), func=funcs, dof=1) idx = out.obs_names.astype(int).to_numpy() mean[idx] = np.asarray(out.layers["mean"]) if need_var: