Skip to content
Open
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
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,9 @@ readme = { file = "README.md", content-type = "text/markdown" }
dynamic = [ "version" ]

dependencies = [
"anndata>=0.10.0",
"anndata>=0.12.14",
"scanpy>=1.10.0",
"scverse-misc[settings]>=0.1.3",
"numpy>=1.17.0",
"scipy>=1.4",
"pandas",
Expand Down
3 changes: 3 additions & 0 deletions src/rapids_singlecell/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@

import cuml.internals.logger as logger

# Import settings before the public modules which consume them.
from ._settings import Preset, settings # isort: skip

from . import dcg, get, gr, pp, ptg, tl
from ._version import __version__

Expand Down
173 changes: 173 additions & 0 deletions src/rapids_singlecell/_settings.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,173 @@
from __future__ import annotations

import enum
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal, NamedTuple, cast

from scverse_misc import Settings as BaseSettings

if TYPE_CHECKING:
from collections.abc import Generator


type HVGFlavor = Literal[
"seurat",
"cell_ranger",
"seurat_v3",
"seurat_v3_paper",
"pearson_residuals",
"poisson_gene_selection",
]
type DETest = Literal[
"logreg", "t-test", "t-test_overestim_var", "wilcoxon", "wilcoxon_binned"
]


class HVGPreset(NamedTuple):
flavor: HVGFlavor
return_df: bool


class BasicEmbeddingPreset(NamedTuple):
key_added: str | None


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


class ScalePreset(NamedTuple):
zero_center: bool | None


class ScoreGenesPreset(NamedTuple):
ctrl_as_ref: bool


class Preset(enum.StrEnum):
"""Presets for :attr:`rapids_singlecell.settings.preset`.

See properties below for details.
"""

ScanpyV1 = "scanpy-v1"
"""Scanpy 1.*’s default settings."""

ScanpyV2Preview = "scanpy-v2-preview"
"""Scanpy 2.*’s feature default settings. (Preview: subject to change!)"""

@property
def highly_variable_genes(self) -> HVGPreset:
return {
Preset.ScanpyV1: HVGPreset(flavor="seurat", return_df=False),
Preset.ScanpyV2Preview: HVGPreset(flavor="seurat_v3_paper", return_df=True),
}[self]

@property
def pca(self) -> BasicEmbeddingPreset:
return self._embedding("pca")

@property
def umap(self) -> BasicEmbeddingPreset:
return self._embedding("umap")

@property
def tsne(self) -> BasicEmbeddingPreset:
return self._embedding("tsne")

@property
def diffmap(self) -> BasicEmbeddingPreset:
return self._embedding("diffmap")

@property
def draw_graph(self) -> BasicEmbeddingPreset:
return BasicEmbeddingPreset(
key_added=None if self is Preset.ScanpyV1 else "graph_{layout}"
)

@property
def rank_genes_groups(self) -> RankGenesGroupsPreset:
return {
Preset.ScanpyV1: RankGenesGroupsPreset(
method="t-test", mask_var=None, mean_in_log_space=True
),
Preset.ScanpyV2Preview: RankGenesGroupsPreset(
method="wilcoxon", mask_var=None, mean_in_log_space=False
),
}[self]

@property
def scale(self) -> ScalePreset:
return ScalePreset(zero_center=True if self is Preset.ScanpyV1 else None)

@property
def score_genes(self) -> ScoreGenesPreset:
return ScoreGenesPreset(ctrl_as_ref=self is Preset.ScanpyV1)

def _embedding(self, name: str) -> BasicEmbeddingPreset:
return BasicEmbeddingPreset(key_added=None if self is Preset.ScanpyV1 else name)

@contextmanager
def override(self, preset: Preset) -> Generator[Preset, None, None]:
"""Temporarily override :attr:`rapids_singlecell.settings.preset`."""
with settings.override(preset=preset):
yield self


class Settings(BaseSettings):
"""Validated global settings for rapids-singlecell."""

preset: Preset = Preset.ScanpyV1
"""Preset to use."""

N_PCS: int = 50
"""Default number of principal components to use."""


settings = Settings()


@dataclass(frozen=True)
class Default:
"""Marker for a function default resolved from :data:`settings`."""

preset: tuple[str, str] | None = None
repr: str | None = None

def __post_init__(self) -> None:
if self.preset is not None and self.repr is not None:
raise TypeError("Cannot provide both preset and repr.")

def resolve(self) -> object:
if self.preset is None:
raise TypeError("A default without a preset cannot be resolved.")
return self._get_value(settings.preset)

def _get_value(self, preset: Preset) -> object:
if self.preset is None:
raise TypeError("A default without a preset has no preset value.")
section, field = self.preset
return getattr(getattr(preset, section), field)

def __repr__(self) -> str:
if self.preset is None:
return self.repr or "default"
value = self.resolve()
suffix = (
" – changes in 2.0"
if settings.preset is Preset.ScanpyV1
and value != self._get_value(Preset.ScanpyV2Preview)
else ""
)
return f"{value!r} (settings.preset={str(settings.preset)!r}{suffix})"


def resolve_default[T](value: T | Default) -> T:
"""Resolve a :class:`Default` marker against the active preset."""
return cast("T", value.resolve()) if isinstance(value, Default) else value


__all__ = ["Preset", "settings"]
9 changes: 9 additions & 0 deletions src/rapids_singlecell/get/_aggregated.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

from rapids_singlecell._compat import DaskArray, _meta_dense
from rapids_singlecell._cuda import _aggr_cuda
from rapids_singlecell._settings import Preset, settings
from rapids_singlecell.get import _check_mask
from rapids_singlecell.preprocessing._utils import _check_gpu_X

Expand Down Expand Up @@ -574,6 +575,14 @@ def aggregate(

Note that this filters out any combination of groups that wasn't present in the original data.
"""
if settings.preset is Preset.ScanpyV2Preview and any(
value is not None for value in (axis, layer, obsm, varm)
):
msg = "`acc` will replace `layer`, `obsm`, and `varm` arguments in Scanpy 2."
if axis is not None:
msg += " `axis` is no longer necessary because it is inferred from `by`."
warnings.warn(msg, FutureWarning, stacklevel=2)

if axis is None:
axis = 1 if varm else 0
axis, axis_name = _resolve_axis(axis)
Expand Down
5 changes: 4 additions & 1 deletion src/rapids_singlecell/preprocessing/_hvg/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import numpy as np

from rapids_singlecell._settings import Default, resolve_default
from rapids_singlecell.preprocessing._utils import _sanitize_column

from ._cutoffs import _Cutoffs
Expand Down Expand Up @@ -37,7 +38,7 @@ def highly_variable_genes(
min_disp: float = 0.5,
max_disp: float = np.inf,
n_top_genes: int = None,
flavor: flavors = "seurat",
flavor: flavors | Default = Default(("highly_variable_genes", "flavor")),
n_bins: int = 20,
span: float = 0.3,
check_values: bool = True,
Expand Down Expand Up @@ -145,6 +146,8 @@ def highly_variable_genes(
`highly_variable_intersection` : bool
If batch_key is given, this denotes the genes that are highly variable in all batches
"""
flavor = resolve_default(flavor)

if batch_key is not None:
_sanitize_column(adata, batch_key)

Expand Down
Loading
Loading