Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/release-notes/4333.feat.md
Original file line number Diff line number Diff line change
@@ -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`
8 changes: 4 additions & 4 deletions docs/tutorials/basics/clustering-2017.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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)"
]
},
Expand Down Expand Up @@ -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)"
]
},
Expand Down Expand Up @@ -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)"
]
Expand Down Expand Up @@ -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",
Expand Down
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -124,10 +124,10 @@ 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",
"sphinx-autodoc-typehints>=3.13.5",
"sphinx-book-theme>=1.1",
"sphinx-copybutton",
"sphinx-design",
Expand Down
29 changes: 28 additions & 1 deletion src/scanpy/_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 2 additions & 3 deletions src/scanpy/_settings/presets.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,6 @@ class BasicEmbeddingPreset(NamedTuple):

class RankGenesGroupsPreset(NamedTuple):
method: DETest
mask_var: str | None
mean_in_log_space: bool


Expand Down Expand Up @@ -245,10 +244,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", mean_in_log_space=True
),
Preset.ScanpyV2Preview: RankGenesGroupsPreset(
method="wilcoxon", mask_var=None, mean_in_log_space=False
method="wilcoxon", mean_in_log_space=False
),
}

Expand Down
9 changes: 5 additions & 4 deletions src/scanpy/_utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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"]) -> 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:
Expand Down
20 changes: 12 additions & 8 deletions src/scanpy/experimental/pp/_normalization.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

import numpy as np
from anndata import AnnData
from scverse_misc import Deprecation, deprecated_arg

from ... import logging as logg
from ... import settings
Expand All @@ -26,6 +27,7 @@
doc_pca_chunk,
)
from ...get import _check_mask, _get_arr, _set_obs_rep
from ...get.get import _mask_arg, _mask_hvg
from ...preprocessing._docs import doc_mask_var
from ...preprocessing._pca import pca

Expand All @@ -34,6 +36,7 @@
from typing import Any

from ..._utils.random import RNGLike, SeedLike
from ...get.get import Mask


def _pearson_residuals(
Expand Down Expand Up @@ -163,11 +166,13 @@ 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)
# `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,
*,
Expand All @@ -176,11 +181,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`.

Expand All @@ -196,7 +201,7 @@ def normalize_pearson_residuals_pca(
{adata}
{dist_params}
{pca_chunk}
{mask_var}
{mask}
{check_values}
{inplace}

Expand Down Expand Up @@ -227,9 +232,8 @@ 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_hvg(adata, _mask_arg(mask, mask_var, dim="var"))
mask_var = _check_mask(adata, mask, "var")

if mask_var is not None:
adata_sub = adata[:, mask_var].copy()
Expand Down
8 changes: 4 additions & 4 deletions src/scanpy/get/_aggregated.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down
71 changes: 69 additions & 2 deletions src/scanpy/get/get.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,8 @@
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
Expand All @@ -38,6 +39,9 @@
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
# --------------------------------------------------------------------------------
Expand Down Expand Up @@ -607,6 +611,37 @@ 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):
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,
Expand All @@ -622,7 +657,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
Expand Down Expand Up @@ -815,6 +850,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`.

Expand Down
13 changes: 7 additions & 6 deletions src/scanpy/preprocessing/_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

from __future__ import annotations

from .._docs import doc_mask

doc_adata_basic = """\
adata
Annotated data matrix.\
Expand All @@ -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
Expand Down
6 changes: 4 additions & 2 deletions src/scanpy/preprocessing/_highly_variable_genes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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=obs_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(
Expand Down
Loading
Loading