Transforms

class cellarium.ml.transforms.BinomialResample(p_binom_min: float, p_binom_max: float, p_apply: float)[source]

Bases: Module

Binomial resampling of gene counts.

For each count, the parameter to the binomial distribution is independently and uniformly sampled according to the bounding parameters, yielding the parameter matrix p_ng.

\[y_{ng} = Binomial(n=x_{ng}, p=p_{ng})\]
Parameters:
  • p_binom_min (float) – Lower bound on binomial distribution parameter.

  • p_binom_max (float) – Upper bound on binomial distribution parameter.

  • p_apply (float) – Probability of applying transform to each sample.

forward(x_ng: Tensor) → dict[str, Tensor][source]
Parameters:

x_ng (Tensor) – Gene counts.

Returns:

Binomially resampled gene counts.

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.CellFilter(min_count_per_cell: int = 0, min_nonzero_genes_per_cell: int = 0)[source]

Bases: Module

Filter cells from the batch by a minimum quality threshold.

Exactly one of min_count_per_cell or min_nonzero_genes_per_cell may be set to a positive value; setting both is an error.

Any torch.Tensor value in the batch whose first dimension matches the number of input cells (x_ng.shape[0]) is filtered consistently. Non-tensor values and tensors whose first dimension differs from n (e.g. gene-indexed arrays) are left unchanged.

When both thresholds are 0 (default) the transform is a no-op: an empty dict is returned and the pipeline merge leaves the batch unmodified. The transform is also a no-op during predict mode (_predict_mode=True in the batch dict, set automatically by predict()).

Warning

CellFilter is incompatible with ContrastiveMLP. The NT-Xent loss asserts a fixed, world-size-divisible batch size and computes positive-pair offsets from it; a variable per-batch cell count breaks both the assertion and the pair alignment, even on a single device.

Parameters:
  • min_count_per_cell (int) – Minimum total count (x_ng row sum) required for a cell to be retained. 0 disables this filter.

  • min_nonzero_genes_per_cell (int) – Minimum number of genes with nonzero expression required for a cell to be retained, computed as (x_ng > 0).sum(dim=-1). 0 disables this filter.

forward(**kwargs: Tensor) → dict[str, Tensor][source]
Parameters:

kwargs (Tensor) – Full batch dictionary forwarded by CellariumPipeline via call_func_with_batch(). Must contain x_ng.

Returns:

A dictionary of filtered tensors to merge back into the batch. Only tensors whose first dimension equals n are included; all others are omitted so the pipeline merge leaves them unchanged. Returns {} when both thresholds are 0 or when _predict_mode is set.

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.CellariumGPTTrainTokenizer(context_len: int, gene_downsample_fraction: float, min_total_mrna_umis: int, max_total_mrna_umis: int, gene_vocab_sizes: dict[str, int], metadata_vocab_sizes: dict[str, int], ontology_downsample_p: float, ontology_infos_path: str, prefix_len: int | None = None, metadata_prompt_token_list: list[str] | None = None, obs_names_rng: bool = False)[source]

Bases: Module

Tokenizer for the Cellarium GPT model.

Parameters:
  • context_len (int) – Context length.

  • gene_downsample_fraction (float) – Fraction of genes to downsample.

  • min_total_mrna_umis (int) – Minimum total mRNA UMIs.

  • max_total_mrna_umis (int) – Maximum total mRNA UMIs.

  • gene_vocab_sizes (dict[str, int]) – Gene token vocabulary sizes.

  • metadata_vocab_sizes (dict[str, int]) – Metadata token vocabulary sizes.

  • ontology_infos_path (str) – Path to ontology information.

  • prefix_len (int | None) – Prefix length. If None, the prefix length is sampled.

  • metadata_prompt_token_list (list[str] | None) – List of metadata tokens to prompt. If None, the metadata prompt tokens are sampled.

  • obs_names_rng (bool) – Cell IDs are used as random seeds for shuffling gene tokens. If None, gene tokens are shuffled without a random seed.

  • ontology_downsample_p (float)

class cellarium.ml.transforms.CenterPerCell(*args: Any, **kwargs: Any)[source]

Bases: Module

Center each cell by subtracting its mean across genes.

Parameters:
  • args (Any)

  • kwargs (Any)

class cellarium.ml.transforms.Densify(*args: Any, **kwargs: Any)[source]

Bases: Module

Convert a sparse x_ng to a dense tensor on the current device.

Use this as the first entry in transforms (GPU transforms) when x_ng arrives as a torch.sparse_csr_tensor (or, on the mps accelerator, a torch.sparse_coo_tensor — see to_torch_sparse_coo()) — for example when no Filter cpu_transform is in the pipeline and the sparse-transfer strategy is still desired (e.g. for statistics models that operate on all genes).

If x_ng is already dense this transform is a no-op.

Parameters:
  • args (Any)

  • kwargs (Any)

forward(x_ng: Tensor, **kwargs: object) → dict[str, Tensor][source]
Parameters:
  • x_ng (Tensor) – Gene counts. May be a sparse (CSR or COO) torch.Tensor or a dense tensor.

  • kwargs (object)

Returns:

A dictionary with the key x_ng containing a dense torch.Tensor.

Raises:

TypeError – If x_ng is a scipy sparse matrix (which should have been converted earlier).

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.DivideByScale(scale_g: Tensor, var_names_g: ndarray, eps: float = 1e-06)[source]

Bases: FilterCompatibilityMixin, Module

Divide gene counts by a scale.

\[y_{ng} = \frac{x_{ng}}{\mathrm{scale}_g + \mathrm{eps}}\]
Parameters:
  • scale_g (Tensor) – A scale for each gene.

  • var_names_g (ndarray) – The variable names schema for the input data validation.

  • eps (float) – A value added to the denominator for numerical stability.

forward(x_ng: Tensor, var_names_g: ndarray) → dict[str, Tensor][source]
Parameters:
  • x_ng (Tensor) – Gene counts.

  • var_names_g (ndarray) – The list of the variable names in the input data. Must be a subset of (or equal to) the var_names_g schema the transform was initialized with, in any order.

Returns:

  • x_ng: The gene counts divided by the scale.

Return type:

A dictionary with the following keys

class cellarium.ml.transforms.Dropout(p_dropout_min, p_dropout_max, p_apply)[source]

Bases: Module

Applies random dropout to gene counts.

For each count, the dropout parameter is independently and uniformly sampled according to the bounding parameters, yielding the parameter matrix p_ng.

\[y_{ng} = x_{ng} * (1 - Bernoulli(p_ng))\]
Parameters:
  • p_dropout_min – Lower bound on dropout parameter.

  • p_dropout_max – Upper bound on dropout parameter.

  • p_apply – Probability of applying transform to each sample.

forward(x_ng: Tensor) → dict[str, Tensor][source]
Parameters:

x_ng (Tensor) – Gene counts.

Returns:

Gene counts with random dropout.

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.Duplicate(enabled=True)[source]

Bases: Module

Duplicates every row of the input tensor, used for contrastive augmentations.

__init__(enabled=True)[source]
Parameters:

enabled – If True, performs duplication; otherwise does nothing. Set False when performing model inference so the transformation pipeline remains consistent with training.

forward(x_ng: Tensor) → dict[str, Tensor][source]
Parameters:

x_ng (Tensor) – Gene counts.

Returns:

Duplicated counts.

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.Filter(filter_list: Sequence[str], ordering: bool = True, allow_missing: bool = False)[source]

Bases: Module

Filter gene counts by a list of features.

When ordering=False, the output columns follow the order genes appear in the input var_names_g:

\[ \begin{align}\begin{aligned}\mathrm{mask}_g = \mathrm{feature}_g \in \mathrm{filter\_list}\\y_{ng} = x_{ng}[:, \mathrm{mask}_g]\end{aligned}\end{align} \]

When ordering=True (default), the output columns follow the order of filter_list:

\[y_{ng} = x_{ng}[:, \sigma(\mathrm{filter\_list})]\]

where \(\sigma\) maps each entry in filter_list to its column index in the input.

Parameters:
  • filter_list (Sequence[str]) – A list of features to filter by.

  • ordering (bool) – If True (default), output columns are ordered to match filter_list. If False, output columns follow the order genes appear in the input var_names_g. Use ordering=True when running inference on data with a different gene ordering than seen during training.

  • allow_missing (bool) – If True, genes in filter_list that are absent from the input are zero-filled in the output. Requires ordering=True. If False (default), all genes in filter_list must be present in the input.

filter(var_names_g: tuple) → ndarray[Any, dtype[int64]] | tuple[ndarray[Any, dtype[int64]], ndarray[Any, dtype[int64]]][source]
Parameters:

var_names_g (tuple) – The list of the variable names in the input data.

Returns:

a 1-D array of source indices in input order.

When ordering=True and allow_missing=False: a 1-D array of source indices ordered to match filter_list.

When ordering=True and allow_missing=True: a tuple (src_indices, out_indices) where src_indices indexes columns in var_names_g and out_indices gives the corresponding destination column in the output (which always has len(filter_list) columns).

Return type:

When ordering=False

forward(x_ng: Tensor | spmatrix, var_names_g: ndarray) → dict[str, Tensor | ndarray][source]

Note

When used with CellariumModule or CellariumPipeline, x_ng and var_names_g keys in the input dictionary will be overwritten with the filtered values.

When x_ng is a scipy.sparse.spmatrix (e.g. when this transform is used as a cpu_transform operating on data from keep_sparse()), a CPU torch.sparse_csr_tensor (e.g. from to_torch_sparse_csr()), or a CPU torch.sparse_coo_tensor (e.g. from to_torch_sparse_coo(), used on the mps accelerator since it has no sparse CSR support), column filtering is performed with scipy (torch has no efficient sparse column indexing) and the result is returned in the same torch sparse layout it arrived in (CSR stays CSR, COO stays COO; plain scipy input becomes CSR). The allow_missing=True path always returns a dense torch.Tensor.

Parameters:
  • x_ng (Tensor | spmatrix) – Gene counts. A dense torch.Tensor, a scipy sparse matrix, or a CPU torch.sparse_csr_tensor or torch.sparse_coo_tensor.

  • var_names_g (ndarray) – The list of the variable names in the input data.

Returns:

  • x_ng: Gene counts filtered (and reordered if ordering=True) to match filter_list. Shape (n, len(filter_list)) when ordering=True, otherwise (n, num_matched).

  • var_names_g: Gene names corresponding to the output columns.

Return type:

A dictionary with the following keys

class cellarium.ml.transforms.GaussianNoise(sigma_min, sigma_max, p_apply)[source]

Bases: Module

Adds Gaussian noise to gene counts.

For each count, Gaussian sigma is independently and uniformly sampled according to the bounding parameters, yielding the sigma matrix sigma_ng.

\[y_{ng} = x_{ng} + N(0, \sigma_{ng})\]
Parameters:
  • sigma_min – Lower bound on Gaussian sigma parameter.

  • sigma_max – Upper bound on Gaussian sigma parameter.

  • p_apply – Probability of applying transform to each sample.

forward(x_ng: Tensor) → dict[str, Tensor][source]
Parameters:

x_ng (Tensor) – Gene counts (log-transformed).

Returns:

Gene counts with added Gaussian noise.

Return type:

dict[str, Tensor]

class cellarium.ml.transforms.Log1p(*args: Any, **kwargs: Any)[source]

Bases: Module

Log1p transform gene counts.

\[y_{ng} = \log(1 + x_{ng})\]
Parameters:
  • args (Any)

  • kwargs (Any)

forward(x_ng: Tensor) → dict[str, Tensor][source]

Note

When used with CellariumModule or CellariumPipeline, x_ng key in the input dictionary will be overwritten with the log1p transformed values.

Parameters:

x_ng (Tensor) – Gene counts.

Returns:

  • x_ng: The log1p transformed gene counts.

Return type:

A dictionary with the following keys

class cellarium.ml.transforms.NormalizeTotal(target_count: int = 10000, eps: float = 1e-06)[source]

Bases: Module

Normalize total gene counts per cell to target count.

\[ \begin{align}\begin{aligned}\mathrm{total\_mrna\_umis}_n = \sum_{g=1}^G x_{ng}\\y_{ng} = \frac{\mathrm{target\_count} \times x_{ng}}{\mathrm{total\_mrna\_umis}_n + \mathrm{eps}}\end{aligned}\end{align} \]
Parameters:
  • target_count (int) – Target gene epxression count.

  • eps (float) – A value added to the denominator for numerical stability.

forward(x_ng: Tensor, total_mrna_umis_n: Tensor | None = None) → dict[str, Tensor][source]

Note

When used with CellariumModule or CellariumPipeline, x_ng key in the input dictionary will be overwritten with the normalized values.

Parameters:
  • x_ng (Tensor) – Gene counts.

  • total_mrna_umis_n (Tensor | None) – Total mRNA UMI counts per cell. If None, it is computed from x_ng.

Returns:

  • x_ng: The gene counts normalized to target count.

Return type:

A dictionary with the following keys

class cellarium.ml.transforms.PFlogPF(target_count: int = 1, eps: float = 1e-06)[source]

Bases: Module

PFlog1pPF / shifted-CLR-style normalization [1] is:

NormalizeTotal + Log1p + CenterPerCell

This is a convenience wrapper for that sequence of transforms, but it provides no additional functionality.

References: [1] Booeshaghi, Hallgrimsdottir, Galvez-Merchan, Pachter. Depth normalization for single-cell genomics count data.

Parameters:
  • target_count (int)

  • eps (float)

class cellarium.ml.transforms.ZScore(mean_g: Tensor, std_g: Tensor, var_names_g: ndarray, eps: float = 1e-06)[source]

Bases: FilterCompatibilityMixin, Module

ZScore gene counts with mean and standard deviation.

\[y_{ng} = \frac{x_{ng} - \mathrm{mean}_g}{\mathrm{std}_g + \mathrm{eps}}\]
Parameters:
  • mean_g (Tensor) – Means for each gene.

  • std_g (Tensor) – Standard deviations for each gene.

  • var_names_g (ndarray) – The variable names schema for the input data validation.

  • eps (float) – A value added to the denominator for numerical stability.

forward(x_ng: Tensor, var_names_g: ndarray) → dict[str, Tensor][source]

Note

When used with CellariumModule or CellariumPipeline, x_ng key in the input dictionary will be overwritten with the z-scored values.

Parameters:
  • x_ng (Tensor) – Gene counts.

  • var_names_g (ndarray) – The list of the variable names in the input data. Must be a subset of (or equal to) the var_names_g schema the transform was initialized with, in any order.

Returns:

  • x_ng: The z-scored gene counts.

Return type:

A dictionary with the following keys