Similarities between <i>Ciona</i> Dorsal Motor Ganglion and Vertebrate Cerebellum: Did a Chordate Ancestor Already Show D/V Subdivision within a Hindbrain Precursor?
The 1 match
- [1] § Materials and Methods › Single-cell RNA sequencing analysis ↔ src/scvi/module/_vae.py, lines 750–893 · score 0.69 · linearly decoded variational, auto encoder, LDVAE, gene expression, space, mapped
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 893 lines · 34 KB · BSD-3-Clause · 1 match
- from __future__ import annotations
- import logging
- import warnings
- from typing import TYPE_CHECKING
- import numpy as np
- import torch
- from torch.nn.functional import one_hot
- from scvi import REGISTRY_KEYS, settings
- from scvi.data._constants import ADATA_MINIFY_TYPE
- from scvi.distributions._utils import _needs_cpu_detour
- from scvi.module._constants import MODULE_KEYS
- from scvi.module.base import (
- BaseMinifiedModeModuleClass,
- EmbeddingModuleMixin,
- LossOutput,
- auto_move_data,
- )
- from scvi.utils import unsupported_if_adata_minified
- if TYPE_CHECKING:
- from collections.abc import Callable
- from typing import Literal
- from torch.distributions import Distribution
- logger = logging.getLogger(__name__)
- class VAE(EmbeddingModuleMixin, BaseMinifiedModeModuleClass):
- """Variational auto-encoder :cite:p:`Lopez18`.
- Parameters
- ----------
- n_input
- Number of input features.
- n_batch
- Number of batches. If ``0``, no batch correction is performed.
- n_labels
- Number of labels.
- n_hidden
- Number of nodes per hidden layer. Passed into :class:`~scvi.nn.Encoder` and
- :class:`~scvi.nn.DecoderSCVI`.
- n_latent
- Dimensionality of the latent space.
- n_layers
- Number of hidden layers. Passed into :class:`~scvi.nn.Encoder` and
- :class:`~scvi.nn.DecoderSCVI`.
- n_continuous_cov
- Number of continuous covariates.
- n_cats_per_cov
- A list of integers containing the number of categories for each categorical covariate.
- dropout_rate
- Dropout rate. Passed into :class:`~scvi.nn.Encoder` but not :class:`~scvi.nn.DecoderSCVI`.
- dispersion
- Flexibility of the dispersion parameter when ``gene_likelihood`` is either ``"nb"`` or
- ``"zinb"``. One of the following:
- * ``"gene"``: parameter is constant per gene across cells.
- * ``"gene-batch"``: parameter is constant per gene per batch.
- * ``"gene-label"``: parameter is constant per gene per label.
- * ``"gene-cell"``: parameter is constant per gene per cell.
- log_variational
- If ``True``, use :func:`~torch.log1p` on input data before encoding for numerical stability
- (not normalization).
- gene_likelihood
- Distribution to use for reconstruction in the generative process. One of the following:
- * ``"nb"``: :class:`~scvi.distributions.NegativeBinomial`.
- * ``"zinb"``: :class:`~scvi.distributions.ZeroInflatedNegativeBinomial`.
- * ``"poisson"``: :class:`~scvi.distributions.Poisson`.
- * ``"normal"``: :class:`~torch.distributions.normal.Normal`.
- latent_distribution
- Distribution to use for the latent space. One of the following:
- * ``"normal"``: isotropic normal.
- * ``"ln"``: logistic normal with normal params N(0, 1).
- encode_covariates
- If ``True``, covariates are concatenated to gene expression prior to passing through
- the encoder(s). Else, only the gene expression is used.
- deeply_inject_covariates
- If ``True`` and ``n_layers > 1``, covariates are concatenated to the outputs of hidden
- layers in the encoder(s) (if ``encoder_covariates`` is ``True``) and the decoder prior to
- passing through the next layer.
- batch_representation
- Method for encoding batch information. One of the following:
- * ``"one-hot"``: represent batches with one-hot encodings.
- * ``"embedding"``: represent batches with continuously-valued embeddings using
- :class:`~scvi.nn.Embedding`.
- Note that batch representations are only passed into the encoder(s) if
- ``encode_covariates`` is ``True``.
- use_batch_norm
- Specifies where to use :class:`~torch.nn.BatchNorm1d` in the model. One of the following:
- * ``"none"``: don't use batch norm in either encoder(s) or decoder.
- * ``"encoder"``: use batch norm only in the encoder(s).
- * ``"decoder"``: use batch norm only in the decoder.
- * ``"both"``: use batch norm in both encoder(s) and decoder.
- Note: if ``use_layer_norm`` is also specified, both will be applied (first
- :class:`~torch.nn.BatchNorm1d`, then :class:`~torch.nn.LayerNorm`).
- use_layer_norm
- Specifies where to use :class:`~torch.nn.LayerNorm` in the model. One of the following:
- * ``"none"``: don't use layer norm in either encoder(s) or decoder.
- * ``"encoder"``: use layer norm only in the encoder(s).
- * ``"decoder"``: use layer norm only in the decoder.
- * ``"both"``: use layer norm in both encoder(s) and decoder.
- Note: if ``use_batch_norm`` is also specified, both will be applied (first
- :class:`~torch.nn.BatchNorm1d`, then :class:`~torch.nn.LayerNorm`).
- use_size_factor_key
- If ``True``, use the :attr:`~anndata.AnnData.obs` column as defined by the
- ``size_factor_key`` parameter in the model's ``setup_anndata`` method as the scaling
- factor in the mean of the conditional distribution. Takes priority over
- ``use_observed_lib_size``.
- use_observed_lib_size
- If ``True``, use the observed library size for RNA as the scaling factor in the mean of the
- conditional distribution.
- extra_payload_autotune
- If ``True``, will return extra matrices in the loss output to be used during autotune
- library_log_means
- :class:`~numpy.ndarray` of shape ``(1, n_batch)`` of means of the log library sizes that
- parameterize the prior on library size if ``use_size_factor_key`` is ``False`` and
- ``use_observed_lib_size`` is ``False``.
- library_log_vars
- :class:`~numpy.ndarray` of shape ``(1, n_batch)`` of variances of the log library sizes
- that parameterize the prior on library size if ``use_size_factor_key`` is ``False`` and
- ``use_observed_lib_size`` is ``False``.
- var_activation
- Callable used to ensure positivity of the variance of the variational distribution. Passed
- into :class:`~scvi.nn.Encoder`. Defaults to :func:`~torch.exp`.
- extra_encoder_kwargs
- Additional keyword arguments passed into :class:`~scvi.nn.Encoder`.
- extra_decoder_kwargs
- Additional keyword arguments passed into :class:`~scvi.nn.DecoderSCVI`.
- batch_embedding_kwargs
- Keyword arguments passed into :class:`~scvi.nn.Embedding` if ``batch_representation`` is
- set to ``"embedding"``.
- """
- def __init__(
- self,
- n_input: int,
- n_batch: int = 0,
- n_labels: int = 0,
- n_hidden: int = 128,
- n_latent: int = 10,
- n_layers: int = 1,
- n_continuous_cov: int = 0,
- n_cats_per_cov: list[int] | None = None,
- dropout_rate: float = 0.1,
- dispersion: Literal["gene", "gene-batch", "gene-label", "gene-cell"] = "gene",
- log_variational: bool = True,
- gene_likelihood: Literal["zinb", "nb", "poisson", "normal"] = "zinb",
- latent_distribution: Literal["normal", "ln"] = "normal",
- encode_covariates: bool = False,
- deeply_inject_covariates: bool = True,
- batch_representation: Literal["one-hot", "embedding"] = "one-hot",
- use_batch_norm: Literal["encoder", "decoder", "none", "both"] = "both",
- use_layer_norm: Literal["encoder", "decoder", "none", "both"] = "none",
- use_size_factor_key: bool = False,
- use_observed_lib_size: bool = True,
- extra_payload_autotune: bool = False,
- library_log_means: np.ndarray | None = None,
- library_log_vars: np.ndarray | None = None,
- var_activation: Callable[[torch.Tensor], torch.Tensor] = None,
- extra_encoder_kwargs: dict | None = None,
- extra_decoder_kwargs: dict | None = None,
- batch_embedding_kwargs: dict | None = None,
- ):
- from scvi.nn import DecoderSCVI, Encoder
- super().__init__()
- self.dispersion = dispersion
- self.n_latent = n_latent
- self.log_variational = log_variational
- self.gene_likelihood = gene_likelihood
- self.n_batch = n_batch
- self.n_input = n_input
- self.n_labels = n_labels
- self.n_hidden = n_hidden
- self.n_layers = n_layers
- self.latent_distribution = latent_distribution
- self.encode_covariates = encode_covariates
- self.use_size_factor_key = use_size_factor_key
- self.use_observed_lib_size = use_size_factor_key or use_observed_lib_size
- self.extra_payload_autotune = extra_payload_autotune
- if not self.use_observed_lib_size:
- if library_log_means is None or library_log_vars is None:
- raise ValueError(
- "If not using observed_lib_size, "
- "must provide library_log_means and library_log_vars."
- )
- self.register_buffer("library_log_means", torch.from_numpy(library_log_means).float())
- self.register_buffer("library_log_vars", torch.from_numpy(library_log_vars).float())
- if self.dispersion == "gene":
- self.px_r = torch.nn.Parameter(torch.randn(n_input))
- elif self.dispersion == "gene-batch":
- self.px_r = torch.nn.Parameter(torch.randn(n_input, n_batch))
- elif self.dispersion == "gene-label":
- self.px_r = torch.nn.Parameter(torch.randn(n_input, n_labels))
- elif self.dispersion == "gene-cell":
- pass
- else:
- raise ValueError(
- "`dispersion` must be one of 'gene', 'gene-batch', 'gene-label', 'gene-cell'."
- )
- self.batch_representation = batch_representation
- if self.batch_representation == "embedding":
- self.init_embedding(REGISTRY_KEYS.BATCH_KEY, n_batch, **(batch_embedding_kwargs or {}))
- batch_dim = self.get_embedding(REGISTRY_KEYS.BATCH_KEY).embedding_dim
- elif self.batch_representation != "one-hot":
- raise ValueError("`batch_representation` must be one of 'one-hot', 'embedding'.")
- use_batch_norm_encoder = use_batch_norm == "encoder" or use_batch_norm == "both"
- use_batch_norm_decoder = use_batch_norm == "decoder" or use_batch_norm == "both"
- use_layer_norm_encoder = use_layer_norm == "encoder" or use_layer_norm == "both"
- use_layer_norm_decoder = use_layer_norm == "decoder" or use_layer_norm == "both"
- n_input_encoder = n_input + n_continuous_cov * encode_covariates
- if self.batch_representation == "embedding":
- n_input_encoder += batch_dim * encode_covariates
- cat_list = list([] if n_cats_per_cov is None else n_cats_per_cov)
- else:
- cat_list = [n_batch] + list([] if n_cats_per_cov is None else n_cats_per_cov)
- encoder_cat_list = cat_list if encode_covariates else None
- _extra_encoder_kwargs = extra_encoder_kwargs or {}
- self.z_encoder = Encoder(
- n_input_encoder,
- n_latent,
- n_cat_list=encoder_cat_list,
- n_layers=n_layers,
- n_hidden=n_hidden,
- dropout_rate=dropout_rate,
- distribution=latent_distribution,
- inject_covariates=deeply_inject_covariates,
- use_batch_norm=use_batch_norm_encoder,
- use_layer_norm=use_layer_norm_encoder,
- var_activation=var_activation,
- return_dist=True,
- **_extra_encoder_kwargs,
- )
- # l encoder goes from n_input-dimensional data to 1-d library size
- self.l_encoder = Encoder(
- n_input_encoder,
- 1,
- n_layers=1,
- n_cat_list=encoder_cat_list,
- n_hidden=n_hidden,
- dropout_rate=dropout_rate,
- inject_covariates=deeply_inject_covariates,
- use_batch_norm=use_batch_norm_encoder,
- use_layer_norm=use_layer_norm_encoder,
- var_activation=var_activation,
- return_dist=True,
- **_extra_encoder_kwargs,
- )
- n_input_decoder = n_latent + n_continuous_cov
- if self.batch_representation == "embedding":
- n_input_decoder += batch_dim
- _extra_decoder_kwargs = extra_decoder_kwargs or {}
- self.decoder = DecoderSCVI(
- n_input_decoder,
- n_input,
- n_cat_list=cat_list,
- n_layers=n_layers,
- n_hidden=n_hidden,
- inject_covariates=deeply_inject_covariates,
- use_batch_norm=use_batch_norm_decoder,
- use_layer_norm=use_layer_norm_decoder,
- scale_activation="softplus" if use_size_factor_key else "softmax",
- **_extra_decoder_kwargs,
- )
- def _get_inference_input(
- self,
- tensors: dict[str, torch.Tensor | None],
- full_forward_pass: bool = False,
- ) -> dict[str, torch.Tensor | None]:
- """Get input tensors for the inference process."""
- if full_forward_pass or self.minified_data_type is None:
- loader = "full_data"
- elif self.minified_data_type in [
- ADATA_MINIFY_TYPE.LATENT_POSTERIOR,
- ADATA_MINIFY_TYPE.LATENT_POSTERIOR_WITH_COUNTS,
- ]:
- loader = "minified_data"
- else:
- raise NotImplementedError(f"Unknown minified-data type: {self.minified_data_type}")
- if loader == "full_data":
- return {
- MODULE_KEYS.X_KEY: tensors[REGISTRY_KEYS.X_KEY],
- MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY],
- MODULE_KEYS.CONT_COVS_KEY: tensors.get(REGISTRY_KEYS.CONT_COVS_KEY, None),
- MODULE_KEYS.CAT_COVS_KEY: tensors.get(REGISTRY_KEYS.CAT_COVS_KEY, None),
- }
- else:
- return {
- MODULE_KEYS.QZM_KEY: tensors[REGISTRY_KEYS.LATENT_QZM_KEY],
- MODULE_KEYS.QZV_KEY: tensors[REGISTRY_KEYS.LATENT_QZV_KEY],
- REGISTRY_KEYS.OBSERVED_LIB_SIZE: tensors[REGISTRY_KEYS.OBSERVED_LIB_SIZE],
- }
- def _get_generative_input(
- self,
- tensors: dict[str, torch.Tensor],
- inference_outputs: dict[str, torch.Tensor | Distribution | None],
- ) -> dict[str, torch.Tensor | None]:
- """Get input tensors for the generative process."""
- size_factor = tensors.get(REGISTRY_KEYS.SIZE_FACTOR_KEY, None)
- if size_factor is not None:
- size_factor = torch.log(size_factor)
- return {
- MODULE_KEYS.Z_KEY: inference_outputs[MODULE_KEYS.Z_KEY],
- MODULE_KEYS.LIBRARY_KEY: inference_outputs[MODULE_KEYS.LIBRARY_KEY],
- MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY],
- MODULE_KEYS.Y_KEY: tensors[REGISTRY_KEYS.LABELS_KEY],
- MODULE_KEYS.CONT_COVS_KEY: tensors.get(REGISTRY_KEYS.CONT_COVS_KEY, None),
- MODULE_KEYS.CAT_COVS_KEY: tensors.get(REGISTRY_KEYS.CAT_COVS_KEY, None),
- MODULE_KEYS.SIZE_FACTOR_KEY: size_factor,
- }
- def _compute_local_library_params(
- self,
- batch_index: torch.Tensor,
- ) -> tuple[torch.Tensor, torch.Tensor]:
- """Computes local library parameters.
- Compute two tensors of shape (batch_index.shape[0], 1) where each
- element corresponds to the mean and variances, respectively, of the
- log library sizes in the batch the cell corresponds to.
- """
- from torch.nn.functional import linear
- n_batch = self.library_log_means.shape[1]
- local_library_log_means = linear(
- one_hot(batch_index.squeeze(-1), n_batch).float(), self.library_log_means
- )
- local_library_log_vars = linear(
- one_hot(batch_index.squeeze(-1), n_batch).float(), self.library_log_vars
- )
- return local_library_log_means, local_library_log_vars
- @auto_move_data
- def _regular_inference(
- self,
- x: torch.Tensor,
- batch_index: torch.Tensor,
- cont_covs: torch.Tensor | None = None,
- cat_covs: torch.Tensor | None = None,
- n_samples: int = 1,
- ) -> dict[str, torch.Tensor | Distribution | None]:
- """Run the regular inference process."""
- x_ = x
- if self.use_observed_lib_size:
- library = torch.log(x.sum(1)).unsqueeze(1)
- if self.log_variational:
- x_ = torch.log1p(x_)
- if cont_covs is not None and self.encode_covariates:
- encoder_input = torch.cat((x_, cont_covs), dim=-1)
- else:
- encoder_input = x_
- if cat_covs is not None and self.encode_covariates:
- categorical_input = torch.split(cat_covs, 1, dim=1)
- else:
- categorical_input = ()
- if self.batch_representation == "embedding" and self.encode_covariates:
- batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index)
- encoder_input = torch.cat([encoder_input, batch_rep], dim=-1)
- qz, z = self.z_encoder(encoder_input, *categorical_input)
- else:
- qz, z = self.z_encoder(encoder_input, batch_index, *categorical_input)
- ql = None
- if not self.use_observed_lib_size:
- if self.batch_representation == "embedding":
- ql, library_encoded = self.l_encoder(encoder_input, *categorical_input)
- else:
- ql, library_encoded = self.l_encoder(
- encoder_input, batch_index, *categorical_input
- )
- library = library_encoded
- if n_samples > 1:
- untran_z = qz.sample((n_samples,))
- z = self.z_encoder.z_transformation(untran_z)
- if self.use_observed_lib_size:
- library = library.unsqueeze(0).expand(
- (n_samples, library.size(0), library.size(1))
- )
- else:
- library = ql.sample((n_samples,))
- return {
- MODULE_KEYS.Z_KEY: z,
- MODULE_KEYS.QZ_KEY: qz,
- MODULE_KEYS.QL_KEY: ql,
- MODULE_KEYS.LIBRARY_KEY: library,
- }
- @auto_move_data
- def _cached_inference(
- self,
- qzm: torch.Tensor,
- qzv: torch.Tensor,
- observed_lib_size: torch.Tensor,
- n_samples: int = 1,
- ) -> dict[str, torch.Tensor | None]:
- """Run the cached inference process."""
- from torch.distributions import Normal
- qz = Normal(qzm, qzv.sqrt())
- # use dist.sample() rather than rsample because we aren't optimizing the z here
- untran_z = qz.sample() if n_samples == 1 else qz.sample((n_samples,))
- z = self.z_encoder.z_transformation(untran_z)
- library = torch.log(observed_lib_size)
- if n_samples > 1:
- library = library.unsqueeze(0).expand((n_samples, library.size(0), library.size(1)))
- return {
- MODULE_KEYS.Z_KEY: z,
- MODULE_KEYS.QZ_KEY: qz,
- MODULE_KEYS.QL_KEY: None,
- MODULE_KEYS.LIBRARY_KEY: library,
- }
- @auto_move_data
- def generative(
- self,
- z: torch.Tensor,
- library: torch.Tensor,
- batch_index: torch.Tensor,
- cont_covs: torch.Tensor | None = None,
- cat_covs: torch.Tensor | None = None,
- size_factor: torch.Tensor | None = None,
- y: torch.Tensor | None = None,
- transform_batch: torch.Tensor | None = None,
- ) -> dict[str, Distribution | None]:
- """Run the generative process."""
- from torch.nn.functional import linear
- from scvi.distributions import (
- NegativeBinomial,
- Normal,
- Poisson,
- ZeroInflatedNegativeBinomial,
- )
- # TODO: refactor forward function to not rely on y
- # Likelihood distribution
- if cont_covs is None:
- decoder_input = z
- elif z.dim() != cont_covs.dim():
- decoder_input = torch.cat(
- [z, cont_covs.unsqueeze(0).expand(z.size(0), -1, -1)], dim=-1
- )
- else:
- decoder_input = torch.cat([z, cont_covs], dim=-1)
- if cat_covs is not None:
- categorical_input = torch.split(cat_covs, 1, dim=1)
- else:
- categorical_input = ()
- if transform_batch is not None:
- batch_index = torch.ones_like(batch_index) * transform_batch
- if not self.use_size_factor_key:
- size_factor = library
- if self.batch_representation == "embedding":
- batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index)
- decoder_input = torch.cat([decoder_input, batch_rep], dim=-1)
- px_scale, px_r, px_rate, px_dropout = self.decoder(
- self.dispersion,
- decoder_input,
- size_factor,
- *categorical_input,
- y,
- )
- else:
- px_scale, px_r, px_rate, px_dropout = self.decoder(
- self.dispersion,
- decoder_input,
- size_factor,
- batch_index,
- *categorical_input,
- y,
- )
- if self.dispersion == "gene-label":
- px_r = linear(
- one_hot(y.squeeze(-1), self.n_labels).float(), self.px_r
- ) # px_r gets transposed - last dimension is nb genes
- elif self.dispersion == "gene-batch":
- px_r = linear(one_hot(batch_index.squeeze(-1), self.n_batch).float(), self.px_r)
- elif self.dispersion == "gene":
- px_r = self.px_r
- px_r = torch.exp(px_r)
- if self.gene_likelihood == "zinb":
- px = ZeroInflatedNegativeBinomial(
- mu=px_rate,
- theta=px_r,
- zi_logits=px_dropout,
- scale=px_scale,
- )
- elif self.gene_likelihood == "nb":
- px = NegativeBinomial(mu=px_rate, theta=px_r, scale=px_scale)
- elif self.gene_likelihood == "poisson":
- px = Poisson(rate=px_rate, scale=px_scale)
- elif self.gene_likelihood == "normal":
- px = Normal(px_rate, px_r, normal_mu=px_scale)
- # Priors
- if self.use_observed_lib_size:
- pl = None
- else:
- (
- local_library_log_means,
- local_library_log_vars,
- ) = self._compute_local_library_params(batch_index)
- pl = Normal(local_library_log_means, local_library_log_vars.sqrt())
- pz = Normal(torch.zeros_like(z), torch.ones_like(z))
- return {
- MODULE_KEYS.PX_KEY: px,
- MODULE_KEYS.PL_KEY: pl,
- MODULE_KEYS.PZ_KEY: pz,
- }
- @unsupported_if_adata_minified
- def loss(
- self,
- tensors: dict[str, torch.Tensor],
- inference_outputs: dict[str, torch.Tensor | Distribution | None],
- generative_outputs: dict[str, Distribution | None],
- kl_weight: torch.Tensor | float = 1.0,
- ) -> LossOutput:
- """Compute the loss."""
- from torch.distributions import kl_divergence
- x = tensors[REGISTRY_KEYS.X_KEY]
- kl_divergence_z = kl_divergence(
- inference_outputs[MODULE_KEYS.QZ_KEY], generative_outputs[MODULE_KEYS.PZ_KEY]
- ).sum(dim=-1)
- if not self.use_observed_lib_size:
- kl_divergence_l = kl_divergence(
- inference_outputs[MODULE_KEYS.QL_KEY], generative_outputs[MODULE_KEYS.PL_KEY]
- ).sum(dim=1)
- else:
- kl_divergence_l = torch.zeros_like(kl_divergence_z)
- reconst_loss = -generative_outputs[MODULE_KEYS.PX_KEY].log_prob(x).sum(-1)
- kl_local_for_warmup = kl_divergence_z
- kl_local_no_warmup = kl_divergence_l
- weighted_kl_local = kl_weight * kl_local_for_warmup + kl_local_no_warmup
- loss = torch.mean(reconst_loss + weighted_kl_local)
- # a payload to be used during autotune
- if self.extra_payload_autotune:
- extra_metrics_payload = {
- "z": inference_outputs["z"],
- "batch": tensors[REGISTRY_KEYS.BATCH_KEY],
- "labels": tensors[REGISTRY_KEYS.LABELS_KEY],
- }
- else:
- extra_metrics_payload = {}
- return LossOutput(
- loss=loss,
- reconstruction_loss=reconst_loss,
- kl_local={
- MODULE_KEYS.KL_L_KEY: kl_divergence_l,
- MODULE_KEYS.KL_Z_KEY: kl_divergence_z,
- },
- extra_metrics=extra_metrics_payload,
- )
- @torch.inference_mode()
- def sample(
- self,
- tensors: dict[str, torch.Tensor],
- n_samples: int = 1,
- max_poisson_rate: float = 1e8,
- generative_kwargs: dict | None = None,
- ) -> torch.Tensor:
- r"""Generate predictive samples from the posterior predictive distribution.
- The posterior predictive distribution is denoted as :math:`p(\hat{x} \mid x)`, where
- :math:`x` is the input data and :math:`\hat{x}` is the sampled data.
- We sample from this distribution by first sampling ``n_samples`` times from the posterior
- distribution :math:`q(z \mid x)` for a given observation, and then sampling from the
- likelihood :math:`p(\hat{x} \mid z)` for each of these.
- Parameters
- ----------
- tensors
- Dictionary of tensors passed into ``VAE.forward``.
- n_samples
- Number of Monte Carlo samples to draw from the distribution for each observation.
- max_poisson_rate
- The maximum value to which to clip the ``rate`` parameter of
- :class:`~scvi.distributions.Poisson`. Avoids numerical sampling issues when the
- parameter is very large due to the variance of the distribution.
- generative_kwargs
- Keyword args for ``generative()`` in fwd pass
- Returns
- -------
- Tensor on CPU with shape ``(n_obs, n_vars)`` if ``n_samples == 1``, else
- ``(n_obs, n_vars,)``.
- """
- from scvi.distributions import Poisson
- inference_kwargs = {"n_samples": n_samples}
- _, generative_outputs = self.forward(
- tensors,
- inference_kwargs=inference_kwargs,
- generative_kwargs=generative_kwargs,
- compute_loss=False,
- )
- dist = generative_outputs[MODULE_KEYS.PX_KEY]
- if self.gene_likelihood == "poisson":
- on_mps = self.device.type == "mps"
- rate = torch.clamp(dist.rate, max=max_poisson_rate)
- if _needs_cpu_detour(on_mps, torch.poisson):
- rate = rate.to("cpu")
- dist = Poisson(rate)
- # (n_obs, n_vars) if n_samples == 1, else (n_samples, n_obs, n_vars)
- samples = dist.sample()
- # (n_samples, n_obs, n_vars) -> (n_obs, n_vars, n_samples)
- samples = torch.permute(samples, (1, 2, 0)) if n_samples > 1 else samples
- return samples.cpu()
- @torch.inference_mode()
- @auto_move_data
- def marginal_ll(
- self,
- tensors: dict[str, torch.Tensor],
- n_mc_samples: int,
- return_mean: bool = False,
- n_mc_samples_per_pass: int = 1,
- ):
- """Compute the marginal log-likelihood of the data under the model.
- Parameters
- ----------
- tensors
- Dictionary of tensors passed into ``VAE.forward``.
- n_mc_samples
- Number of Monte Carlo samples to use for the estimation of the marginal log-likelihood.
- return_mean
- Whether to return the mean of marginal likelihoods over cells.
- n_mc_samples_per_pass
- Number of Monte Carlo samples to use per pass. This is useful to avoid memory issues.
- """
- from torch import logsumexp
- from torch.distributions import Normal
- batch_index = tensors[REGISTRY_KEYS.BATCH_KEY]
- to_sum = []
- if n_mc_samples_per_pass > n_mc_samples:
- warnings.warn(
- "Number of chunks is larger than the total number of samples, setting it to the "
- "number of samples",
- RuntimeWarning,
- stacklevel=settings.warnings_stacklevel,
- )
- n_mc_samples_per_pass = n_mc_samples
- n_passes = int(np.ceil(n_mc_samples / n_mc_samples_per_pass))
- for _ in range(n_passes):
- # Distribution parameters and sampled variables
- inference_outputs, _, losses = self.forward(
- tensors,
- inference_kwargs={"n_samples": n_mc_samples_per_pass},
- get_inference_input_kwargs={"full_forward_pass": True},
- )
- qz = inference_outputs[MODULE_KEYS.QZ_KEY]
- ql = inference_outputs[MODULE_KEYS.QL_KEY]
- z = inference_outputs[MODULE_KEYS.Z_KEY]
- library = inference_outputs[MODULE_KEYS.LIBRARY_KEY]
- # Reconstruction Loss
- reconst_loss = losses.dict_sum(losses.reconstruction_loss)
- # Log-probabilities
- p_z = (
- Normal(torch.zeros_like(qz.loc), torch.ones_like(qz.scale)).log_prob(z).sum(dim=-1)
- )
- p_x_zl = -reconst_loss
- q_z_x = qz.log_prob(z).sum(dim=-1)
- log_prob_sum = p_z + p_x_zl - q_z_x
- if not self.use_observed_lib_size:
- (
- local_library_log_means,
- local_library_log_vars,
- ) = self._compute_local_library_params(batch_index)
- p_l = (
- Normal(local_library_log_means, local_library_log_vars.sqrt())
- .log_prob(library)
- .sum(dim=-1)
- )
- q_l_x = ql.log_prob(library).sum(dim=-1)
- log_prob_sum += p_l - q_l_x
- if n_mc_samples_per_pass == 1:
- log_prob_sum = log_prob_sum.unsqueeze(0)
- to_sum.append(log_prob_sum)
- to_sum = torch.cat(to_sum, dim=0)
- batch_log_lkl = logsumexp(to_sum, dim=0) - np.log(n_mc_samples)
- if return_mean:
- batch_log_lkl = torch.mean(batch_log_lkl).item()
- else:
- batch_log_lkl = batch_log_lkl.cpu()
- return batch_log_lkl
- class LDVAE(VAE):
- """Linear-decoded Variational auto-encoder model.
- Implementation of :cite:p:`Svensson20`.
- This model uses a linear decoder, directly mapping the latent representation
- to gene expression levels. It still uses a deep neural network to encode
- the latent representation.
- Compared to standard VAE, this model is less powerful but can be used to
- inspect which genes contribute to variation in the dataset. It may also be used
- for all scVI tasks, like differential expression, batch correction, imputation, etc.
- However, batch correction may be less powerful as it assumes a linear model.
- Parameters
- ----------
- n_input
- Number of input genes
- n_batch
- Number of batches
- n_labels
- Number of labels
- n_hidden
- Number of nodes per hidden layer (for encoder)
- n_latent
- Dimensionality of the latent space
- n_layers_encoder
- Number of hidden layers used for encoder NNs
- dropout_rate
- Dropout rate for neural networks
- dispersion
- One of the following
- * ``'gene'`` - dispersion parameter of NB is constant per gene across cells
- * ``'gene-batch'`` - dispersion can differ between different batches
- * ``'gene-label'`` - dispersion can differ between different labels
- * ``'gene-cell'`` - dispersion can differ for every gene in every cell
- log_variational
- Log(data+1) prior to encoding for numerical stability. Not normalization.
- gene_likelihood
- One of
- * ``'nb'`` - Negative binomial distribution
- * ``'zinb'`` - Zero-inflated negative binomial distribution
- * ``'poisson'`` - Poisson distribution
- use_batch_norm
- Bool whether to use batch norm in decoder
- bias
- Bool whether to have bias term in linear decoder
- latent_distribution
- One of
- * ``'normal'`` - Isotropic normal
- * ``'ln'`` - Logistic normal with normal params N(0, 1)
- use_observed_lib_size
- Use observed library size for RNA as scaling factor in mean of conditional distribution.
- **kwargs
- """
- def __init__(
- self,
- n_input: int,
- n_batch: int = 0,
- n_labels: int = 0,
- n_hidden: int = 128,
- n_latent: int = 10,
- n_layers_encoder: int = 1,
- dropout_rate: float = 0.1,
- dispersion: str = "gene",
- log_variational: bool = True,
- gene_likelihood: str = "nb",
- use_batch_norm: bool = True,
- bias: bool = False,
- latent_distribution: str = "normal",
- use_observed_lib_size: bool = False,
- **kwargs,
- ):
- from scvi.nn import Encoder, LinearDecoderSCVI
- super().__init__(
- n_input=n_input,
- n_batch=n_batch,
- n_labels=n_labels,
- n_hidden=n_hidden,
- n_latent=n_latent,
- n_layers=n_layers_encoder,
- dropout_rate=dropout_rate,
- dispersion=dispersion,
- log_variational=log_variational,
- gene_likelihood=gene_likelihood,
- latent_distribution=latent_distribution,
- use_observed_lib_size=use_observed_lib_size,
- **kwargs,
- )
- self.use_batch_norm = use_batch_norm
- self.z_encoder = Encoder(
- n_input,
- n_latent,
- n_layers=n_layers_encoder,
- n_hidden=n_hidden,
- dropout_rate=dropout_rate,
- distribution=latent_distribution,
- use_batch_norm=True,
- use_layer_norm=False,
- return_dist=True,
- )
- self.l_encoder = Encoder(
- n_input,
- 1,
- n_layers=1,
- n_hidden=n_hidden,
- dropout_rate=dropout_rate,
- use_batch_norm=True,
- use_layer_norm=False,
- return_dist=True,
- )
- self.decoder = LinearDecoderSCVI(
- n_latent,
- n_input,
- n_cat_list=[n_batch],
- use_batch_norm=use_batch_norm,
- use_layer_norm=False,
- bias=bias,
- )
- @torch.inference_mode()
- def get_loadings(self) -> np.ndarray:
- """Extract per-gene weights in the linear decoder."""
- # This is BW, where B is diag(b) batch norm, W is the weight matrix
- if self.use_batch_norm is True:
- w = self.decoder.factor_regressor.fc_layers[0][0].weight
- bn = self.decoder.factor_regressor.fc_layers[0][1]
- sigma = torch.sqrt(bn.running_var + bn.eps)
- gamma = bn.weight
- b = gamma / sigma
- b_identity = torch.diag(b)
- loadings = torch.matmul(b_identity, w)
- else:
- loadings = self.decoder.factor_regressor.fc_layers[0][0].weight
- loadings = loadings.detach().cpu().numpy()
- if self.n_batch > 1:
- loadings = loadings[:, : -self.n_batch]
- return loadings
_vae.py at commit fa9470f, under BSD-3-Clause · at the source
Overview
- Neuroscience Research Institute, University of California Santa Barbara, Santa Barbara, California 93106
- Department of Molecular, Cellular and Developmental Biology, University of California, Santa Barbara, California 93106
- Life Sciences Centre, Dalhousie University, Halifax, Nova Scotia B3H 1A5, Canada
Abstract
The cerebellum, a major structure of the vertebrate central nervous system, forms the dorsal region of the hindbrain. While evidence from connectomics, neuron classification, and gene expression analyses suggests that the major CNS divisions of forebrain, midbrain, hindbrain, and spinal cord predate the vertebrate/
Reproduced under the paper's license (CC BY), from the paper cited above.
Repository
Its files are read in the Code ↔ Paper reader above, with 1 match between paragraphs and lines of code.
scverse/scvi-tools
fa9470f9b59f47a1df1677eda44161ac14845db3, 24 September 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
345 files
- conftest.py, Python, 62 lines
- docs/
_static/ , JavaScript, 164 linesjs/ custom.js - docs/
conf.py , Python, 343 lines - docs/
extensions/ , Python, 39 linesedit_colab_url.py - docs/
extensions/ , Python, 319 linesfilterlist.py - docs/
extensions/ , Python, 35 linestyped_returns.py - src/
scvi/ , Python, 31 lines__init__.py - src/
scvi/ , Python, 21 lines_constants.py - src/
scvi/ , Python, 231 lines_settings.py - src/
scvi/ , Python, 13 lines_types.py - src/
scvi/ , Python, 34 linesautotune/ __init__.py - src/
scvi/ , Python, 896 linesautotune/ _experiment.py - src/
scvi/ , Python, 181 linesautotune/ _tune.py - src/
scvi/ , Python, 4 linescriticism/ __init__.py - src/
scvi/ , Python, 14 linescriticism/ _constants.py - src/
scvi/ , Python, 226 linescriticism/ _create_criticism_report .py - src/
scvi/ , Python, 481 linescriticism/ _ppc.py - src/
scvi/ , Python, 68 linesdata/ __init__.py - src/
scvi/ , Python, 169 linesdata/ _anntorchdataset.py - src/
scvi/ , Python, 1 linedata/ _built_in_data/ __init__.py - src/
scvi/ , Python, 112 linesdata/ _built_in_data/ _brain_large.py - src/
scvi/ , Python, 62 linesdata/ _built_in_data/ _cellxgene.py - src/
scvi/ , Python, 187 linesdata/ _built_in_data/ _cite_seq.py - src/
scvi/ , Python, 58 linesdata/ _built_in_data/ _cortex.py - src/
scvi/ , Python, 49 linesdata/ _built_in_data/ _csv.py - src/
scvi/ , Python, 171 linesdata/ _built_in_data/ _dataset_10x.py - src/
scvi/ , Python, 47 linesdata/ _built_in_data/ _heartcellatlas.py - src/
scvi/ , Python, 103 linesdata/ _built_in_data/ _loom.py - src/
scvi/ , Python, 121 linesdata/ _built_in_data/ _pbmc.py - src/
scvi/ , Python, 81 linesdata/ _built_in_data/ _smfish.py - src/
scvi/ , Python, 113 linesdata/ _built_in_data/ _synthetic.py - src/
scvi/ , Python, 173 linesdata/ _compat.py - src/
scvi/ , Python, 60 linesdata/ _constants.py - src/
scvi/ , Python, 698 linesdata/ _datasets.py - src/
scvi/ , Python, 116 linesdata/ _download.py - src/
scvi/ , Python, 570 linesdata/ _manager.py - src/
scvi/ , Python, 492 linesdata/ _preprocessing.py - src/
scvi/ , Python, 71 linesdata/ _read.py - src/
scvi/ , Python, 383 linesdata/ _utils.py - src/
scvi/ , Python, 64 linesdata/ fields/ __init__.py - src/
scvi/ , Python, 522 linesdata/ fields/ _arraylike_field.py - src/
scvi/ , Python, 151 linesdata/ fields/ _base_field.py - src/
scvi/ , Python, 293 linesdata/ fields/ _dataframe_field.py - src/
scvi/ , Python, 147 linesdata/ fields/ _layer_field.py - src/
scvi/ , Python, 143 linesdata/ fields/ _mudata.py - src/
scvi/ , Python, 165 linesdata/ fields/ _protein.py - src/
scvi/ , Python, 113 linesdata/ fields/ _scanvi.py - src/
scvi/ , Python, 84 linesdata/ fields/ _uns_field.py - src/
scvi/ , Python, 33 linesdataloaders/ __init__.py - src/
scvi/ , Python, 138 linesdataloaders/ _ann_dataloader.py - src/
scvi/ , Python, 120 linesdataloaders/ _anncollection.py - src/
scvi/ , Python, 95 linesdataloaders/ _concat_dataloader.py - src/
scvi/ , Python, 1,521 linesdataloaders/ _custom_dataloaders.py - src/
scvi/ , Python, 771 linesdataloaders/ _data_splitting.py - src/
scvi/ , Python, 89 linesdataloaders/ _samplers.py - src/
scvi/ , Python, 114 linesdataloaders/ _semi_dataloader.py - src/
scvi/ , Python, 24 linesdistributions/ __init__.py - src/
scvi/ , Python, 153 linesdistributions/ _beta_binomial.py - src/
scvi/ , Python, 41 linesdistributions/ _constraints.py - src/
scvi/ , Python, 201 linesdistributions/ _gamma.py - src/
scvi/ , Python, 316 linesdistributions/ _lognormal.py - src/
scvi/ , Python, 724 linesdistributions/ _negative_binomial.py - src/
scvi/ , Python, 57 linesdistributions/ _normal.py - src/
scvi/ , Python, 94 linesdistributions/ _utils.py - src/
scvi/ , Python, 50 linesexternal/ __init__.py - src/
scvi/ , Python, 4 linesexternal/ cellassign/ __init__.py - src/
scvi/ , Python, 294 linesexternal/ cellassign/ _model.py - src/
scvi/ , Python, 265 linesexternal/ cellassign/ _module.py - src/
scvi/ , Python, 11 linesexternal/ contrastivevi/ __init__.py - src/
scvi/ , Python, 218 linesexternal/ contrastivevi/ _contrastive_data_splitt ing.py - src/
scvi/ , Python, 123 linesexternal/ contrastivevi/ _contrastive_dataloader. py - src/
scvi/ , Python, 860 linesexternal/ contrastivevi/ _model.py - src/
scvi/ , Python, 610 linesexternal/ contrastivevi/ _module.py - src/
scvi/ , Python, 15 linesexternal/ cytovi/ __init__.py - src/
scvi/ , Python, 22 linesexternal/ cytovi/ _constants.py - src/
scvi/ , Python, 1,241 linesexternal/ cytovi/ _model.py - src/
scvi/ , Python, 576 linesexternal/ cytovi/ _module.py - src/
scvi/ , Python, 268 linesexternal/ cytovi/ _plotting.py - src/
scvi/ , Python, 408 linesexternal/ cytovi/ _preprocessing.py - src/
scvi/ , Python, 270 linesexternal/ cytovi/ _read_write.py - src/
scvi/ , Python, 229 linesexternal/ cytovi/ _utils.py - src/
scvi/ , Python, 4 linesexternal/ decipher/ __init__.py - src/
scvi/ , Python, 99 linesexternal/ decipher/ _components.py - src/
scvi/ , Python, 360 linesexternal/ decipher/ _model.py - src/
scvi/ , Python, 183 linesexternal/ decipher/ _module.py - src/
scvi/ , Python, 143 linesexternal/ decipher/ _trainingplan.py - src/
scvi/ , Python, 4 linesexternal/ decipher/ utils/ __init__.py - src/
scvi/ , Python, 109 linesexternal/ decipher/ utils/ _rotate.py - src/
scvi/ , Python, 138 linesexternal/ decipher/ utils/ _trajectory.py - src/
scvi/ , Python, 7 linesexternal/ diagvi/ __init__.py - src/
scvi/ , Python, 270 linesexternal/ diagvi/ _base_components.py - src/
scvi/ , Python, 1,245 linesexternal/ diagvi/ _model.py - src/
scvi/ , Python, 597 linesexternal/ diagvi/ _module.py - src/
scvi/ , Python, 525 linesexternal/ diagvi/ _task.py - src/
scvi/ , Python, 384 linesexternal/ diagvi/ _utils.py - src/
scvi/ , Python, 14 linesexternal/ drvi/ __init__.py - src/
scvi/ , Python, 313 linesexternal/ drvi/ _base_components.py - src/
scvi/ , Python, 10 linesexternal/ drvi/ _constants.py - src/
scvi/ , Python, 132 linesexternal/ drvi/ _distributions.py - src/
scvi/ , Python, 156 linesexternal/ drvi/ _generative_mixin.py - src/
scvi/ , Python, 722 linesexternal/ drvi/ _interpretability_mixin. py - src/
scvi/ , Python, 109 linesexternal/ drvi/ _model.py - src/
scvi/ , Python, 348 linesexternal/ drvi/ _module.py - src/
scvi/ , Python, 48 linesexternal/ drvi/ _trainingplan.py - src/
scvi/ , Python, 244 linesexternal/ drvi/ _utils.py - src/
scvi/ , Python, 4 linesexternal/ gimvi/ __init__.py - src/
scvi/ , Python, 690 linesexternal/ gimvi/ _model.py - src/
scvi/ , Python, 529 linesexternal/ gimvi/ _module.py - src/
scvi/ , Python, 195 linesexternal/ gimvi/ _task.py - src/
scvi/ , Python, 125 linesexternal/ gimvi/ _utils.py - src/
scvi/ , Python, 4 linesexternal/ joint_embedding_scvi/ __init__.py - src/
scvi/ , Python, 157 linesexternal/ joint_embedding_scvi/ _model.py - src/
scvi/ , Python, 222 linesexternal/ joint_embedding_scvi/ _module.py - src/
scvi/ , Python, 159 linesexternal/ joint_embedding_scvi/ _utils.py - src/
scvi/ , Python, 17 linesexternal/ methylvi/ __init__.py - src/
scvi/ , Python, 525 linesexternal/ methylvi/ _base_components.py - src/
scvi/ , Python, 9 linesexternal/ methylvi/ _constants.py - src/
scvi/ , Python, 275 linesexternal/ methylvi/ _methylanvi_model.py - src/
scvi/ , Python, 270 linesexternal/ methylvi/ _methylanvi_module.py - src/
scvi/ , Python, 267 linesexternal/ methylvi/ _methylvi_model.py - src/
scvi/ , Python, 283 linesexternal/ methylvi/ _methylvi_module.py - src/
scvi/ , Python, 68 linesexternal/ methylvi/ _utils.py - src/
scvi/ , Python, 5 linesexternal/ mrvi/ __init__.py - src/
scvi/ , Python, 432 linesexternal/ mrvi/ _components.py - src/
scvi/ , Python, 2,083 linesexternal/ mrvi/ _model.py - src/
scvi/ , Python, 714 linesexternal/ mrvi/ _module.py - src/
scvi/ , Python, 103 linesexternal/ mrvi/ _types.py - src/
scvi/ , Python, 25 linesexternal/ mrvi/ _utils.py - src/
scvi/ , Python, 3 linesexternal/ poissonvi/ __init__.py - src/
scvi/ , Python, 439 linesexternal/ poissonvi/ _model.py - src/
scvi/ , Python, 4 linesexternal/ resolvi/ __init__.py - src/
scvi/ , Python, 748 linesexternal/ resolvi/ _model.py - src/
scvi/ , Python, 1,329 linesexternal/ resolvi/ _module.py - src/
scvi/ , Python, 566 linesexternal/ resolvi/ _utils.py - src/
scvi/ , Python, 4 linesexternal/ scar/ __init__.py - src/
scvi/ , Python, 385 linesexternal/ scar/ _model.py - src/
scvi/ , Python, 384 linesexternal/ scar/ _module.py - src/
scvi/ , Python, 4 linesexternal/ scbasset/ __init__.py - src/
scvi/ , Python, 446 linesexternal/ scbasset/ _model.py - src/
scvi/ , Python, 398 linesexternal/ scbasset/ _module.py - src/
scvi/ , Python, 13 linesexternal/ scviva/ __init__.py - src/
scvi/ , Python, 262 linesexternal/ scviva/ _components.py - src/
scvi/ , Python, 30 linesexternal/ scviva/ _constants.py - src/
scvi/ , Python, 109 linesexternal/ scviva/ _log_likelihood.py - src/
scvi/ , Python, 1,173 linesexternal/ scviva/ _model.py - src/
scvi/ , Python, 638 linesexternal/ scviva/ _module.py - src/
scvi/ , Python, 11 linesexternal/ scviva/ differential_expression/ __init__.py - src/
scvi/ , Python, 113 linesexternal/ scviva/ differential_expression/ _de_utils.py - src/
scvi/ , Python, 105 linesexternal/ scviva/ differential_expression/ _marker_classifier.py - src/
scvi/ , Python, 233 linesexternal/ scviva/ differential_expression/ _niche_de_core.py - src/
scvi/ , Python, 259 linesexternal/ scviva/ differential_expression/ _results_dataclass.py - src/
scvi/ , Python, 3 linesexternal/ solo/ __init__.py - src/
scvi/ , Python, 484 linesexternal/ solo/ _model.py - src/
scvi/ , Python, 4 linesexternal/ stereoscope/ __init__.py - src/
scvi/ , Python, 394 linesexternal/ stereoscope/ _model.py - src/
scvi/ , Python, 265 linesexternal/ stereoscope/ _module.py - src/
scvi/ , Python, 4 linesexternal/ sysvi/ __init__.py - src/
scvi/ , Python, 221 linesexternal/ sysvi/ _base_components.py - src/
scvi/ , Python, 259 linesexternal/ sysvi/ _model.py - src/
scvi/ , Python, 527 linesexternal/ sysvi/ _module.py - src/
scvi/ , Python, 206 linesexternal/ sysvi/ _priors.py - src/
scvi/ , Python, 4 linesexternal/ tangram/ __init__.py - src/
scvi/ , Python, 355 linesexternal/ tangram/ _model.py - src/
scvi/ , Python, 164 linesexternal/ tangram/ _module.py - src/
scvi/ , Python, 4 linesexternal/ totalanvi/ __init__.py - src/
scvi/ , Python, 540 linesexternal/ totalanvi/ _model.py - src/
scvi/ , Python, 452 linesexternal/ totalanvi/ _module.py - src/
scvi/ , Python, 4 linesexternal/ velovi/ __init__.py - src/
scvi/ , Python, 9 linesexternal/ velovi/ _constants.py - src/
scvi/ , Python, 1,153 linesexternal/ velovi/ _model.py - src/
scvi/ , Python, 656 linesexternal/ velovi/ _module.py - src/
scvi/ , Python, 15 linesexternal/ vivs/ __init__.py - src/
scvi/ , Python, 10 linesexternal/ vivs/ _constants.py - src/
scvi/ , Python, 708 linesexternal/ vivs/ _model.py - src/
scvi/ , Python, 211 linesexternal/ vivs/ _module.py - src/
scvi/ , Python, 137 linesexternal/ vivs/ _plotting.py - src/
scvi/ , Python, 94 linesexternal/ vivs/ _utils.py - src/
scvi/ , Python, 8 lineshub/ __init__.py - src/
scvi/ , Python, 26 lineshub/ _constants.py - src/
scvi/ , Python, 364 lineshub/ _metadata.py - src/
scvi/ , Python, 621 lineshub/ _model.py - src/
scvi/ , Python, 210 lineshub/ _template.py - src/
scvi/ , Python, 57 lineshub/ _url.py - src/
scvi/ , Python, 51 linesmodel/ __init__.py - src/
scvi/ , Python, 272 linesmodel/ _amortizedlda.py - src/
scvi/ , Python, 295 linesmodel/ _autozi.py - src/
scvi/ , Python, 471 linesmodel/ _condscvi.py - src/
scvi/ , Python, 700 linesmodel/ _destvi.py - src/
scvi/ , Python, 192 linesmodel/ _linear_scvi.py - src/
scvi/ , Python, 216 linesmodel/ _mlxscvi.py - src/
scvi/ , Python, 1,282 linesmodel/ _multivi.py - src/
scvi/ , Python, 665 linesmodel/ _peakvi.py - src/
scvi/ , Python, 345 linesmodel/ _scanvi.py - src/
scvi/ , Python, 258 linesmodel/ _scvi.py - src/
scvi/ , Python, 1,557 linesmodel/ _totalvi.py - src/
scvi/ , Python, 427 linesmodel/ _utils.py - src/
scvi/ , Python, 58 linesmodel/ base/ __init__.py - src/
scvi/ , Python, 540 linesmodel/ base/ _archesmixin.py - src/
scvi/ , Python, 2,016 linesmodel/ base/ _base_model.py - src/
scvi/ , Python, 16 linesmodel/ base/ _constants.py - src/
scvi/ , Python, 219 linesmodel/ base/ _de_core.py - src/
scvi/ , Python, 743 linesmodel/ base/ _differential.py - src/
scvi/ , Python, 38 linesmodel/ base/ _embedding_mixin.py - src/
scvi/ , Python, 136 linesmodel/ base/ _log_likelihood.py - src/
scvi/ , Python, 283 linesmodel/ base/ _mlxmixin.py - src/
scvi/ , Python, 718 linesmodel/ base/ _pyromixin.py - src/
scvi/ , Python, 963 linesmodel/ base/ _rnamixin.py - src/
scvi/ , Python, 246 linesmodel/ base/ _save_load.py - src/
scvi/ , Python, 714 linesmodel/ base/ _training_mixin.py - src/
scvi/ , Python, 540 linesmodel/ base/ _vaemixin.py - src/
scvi/ , Python, 17 linesmodel/ utils/ __init__.py - src/
scvi/ , Python, 471 linesmodel/ utils/ _annbatch_de_core.py - src/
scvi/ , Python, 67 linesmodel/ utils/ _minification.py - src/
scvi/ , Python, 48 linesmodule/ __init__.py - src/
scvi/ , Python, 361 linesmodule/ _amortizedlda.py - src/
scvi/ , Python, 413 linesmodule/ _autozivae.py - src/
scvi/ , Python, 79 linesmodule/ _classifier.py - src/
scvi/ , Python, 27 linesmodule/ _constants.py - src/
scvi/ , Python, 392 linesmodule/ _mlxvae.py - src/
scvi/ , Python, 676 linesmodule/ _mrdeconv.py - src/
scvi/ , Python, 1,060 linesmodule/ _multivae.py - src/
scvi/ , Python, 348 linesmodule/ _peakvae.py - src/
scvi/ , Python, 278 linesmodule/ _scanvae.py - src/
scvi/ , Python, 866 linesmodule/ _totalvae.py - src/
scvi/ , Python, 32 linesmodule/ _utils.py - src/
scvi/ , Python, 893 lines, 1 matchmodule/ _vae.py - src/
scvi/ , Python, 382 linesmodule/ _vaec.py - src/
scvi/ , Python, 25 linesmodule/ base/ __init__.py - src/
scvi/ , Python, 595 linesmodule/ base/ _base_module.py - src/
scvi/ , Python, 116 linesmodule/ base/ _decorators.py - src/
scvi/ , Python, 50 linesmodule/ base/ _embedding_mixin.py - src/
scvi/ , Python, 301 linesmodule/ base/ _priors.py - src/
scvi/ , Python, 11 linesmodule/ base/ _pyro.py - src/
scvi/ , Python, 28 linesnn/ __init__.py - src/
scvi/ , Python, 1,099 linesnn/ _base_components.py - src/
scvi/ , Python, 98 linesnn/ _embedding.py - src/
scvi/ , Python, 10 linesnn/ _utils.py - src/
scvi/ , Python, 70 linestrain/ __init__.py - src/
scvi/ , Python, 285 linestrain/ _callbacks.py - src/
scvi/ , Python, 388 linestrain/ _config.py - src/
scvi/ , Python, 17 linestrain/ _constants.py - src/
scvi/ , Python, 124 linestrain/ _logger.py - src/
scvi/ , Python, 91 linestrain/ _metrics.py - src/
scvi/ , Python, 91 linestrain/ _progress.py - src/
scvi/ , Python, 228 linestrain/ _trainer.py - src/
scvi/ , Python, 1,927 linestrain/ _trainingplans.py - src/
scvi/ , Python, 269 linestrain/ _trainrunner.py - src/
scvi/ , Python, 21 linesutils/ __init__.py - src/
scvi/ , Python, 12 linesutils/ _attrdict.py - src/
scvi/ , Python, 18 linesutils/ _decorators.py - src/
scvi/ , Python, 44 linesutils/ _dependencies.py - src/
scvi/ , Python, 258 linesutils/ _docstrings.py - src/
scvi/ , Python, 144 linesutils/ _mlflow.py - src/
scvi/ , Python, 59 linesutils/ _track.py - tests/
__init__.py , Python, 1 line - tests/
autotune/ , Python, 1,181 linestest_experiment.py - tests/
autotune/ , Python, 336 linestest_tune.py - tests/
conftest.py , Python, 168 lines - tests/
criticism/ , Python, 122 linestest_criticism.py - tests/
data/ , Python, 1 line__init__.py - tests/
data/ , Python, 20 linesconftest.py - tests/
data/ , Python, 548 linestest_anndata.py - tests/
data/ , Python, 176 linestest_anntorchdataset.py - tests/
data/ , Python, 112 linestest_built_in_data.py - tests/
data/ , Python, 32 linestest_data_utils.py - tests/
data/ , Python, 47 linestest_dataset10X.py - tests/
data/ , Python, 394 linestest_mudata.py - tests/
data/ , Python, 104 linestest_preprocessing.py - tests/
data/ , Python, 17 linestest_synthetic_iid.py - tests/
data/ , Python, 142 linesutils.py - tests/
dataloaders/ , Python, 1 line__init__.py - tests/
dataloaders/ , Python, 118 linessparse_utils.py - tests/
dataloaders/ , Python, 1,014 linestest_annbatch.py - tests/
dataloaders/ , Python, 1,205 linestest_custom_dataloader.p y - tests/
dataloaders/ , Python, 301 linestest_dataloaders.py - tests/
dataloaders/ , Python, 272 linestest_datasplitter.py - tests/
dataloaders/ , Python, 155 linestest_samplers.py - tests/
distributions/ , Python, 77 linestest_beta_binomial.py - tests/
distributions/ , Python, 21 linestest_constraints.py - tests/
distributions/ , Python, 154 linestest_gamma.py - tests/
distributions/ , Python, 146 linestest_lognormal.py - tests/
distributions/ , Python, 230 linestest_negative_binomial.p y - tests/
external/ , Python, 61 linescellassign/ test_model_cellassign.py - tests/
external/ , Python, 54 linescontrastivevi/ test_contrastive_dataloa ders.py - tests/
external/ , Python, 158 linescontrastivevi/ test_contrastive_dataspl itter.py - tests/
external/ , Python, 458 linescontrastivevi/ test_contrastivevae.py - tests/
external/ , Python, 165 linescontrastivevi/ test_contrastivevi.py - tests/
external/ , Python, 301 linescytovi/ test_cytovi.py - tests/
external/ , Python, 135 linesdecipher/ test_decipher.py - tests/
external/ , Python, 1,680 linesdiagvi/ test_diagvi.py - tests/
external/ , Python, 641 linesdrvi/ test_drvi.py - tests/
external/ , Python, 232 linesgimvi/ test_gimvi.py - tests/
external/ , Python, 151 linesjoint_embedding_scvi/ test_joint_embedding_scv i.py - tests/
external/ , Python, 61 linesjoint_embedding_scvi/ test_joint_embedding_uti ls.py - tests/
external/ , Python, 45 linesmethylvi/ test_methylanvi.py - tests/
external/ , Python, 62 linesmethylvi/ test_methylvi.py - tests/
external/ , Python, 82 linesmrvi/ test_mrvi_components.py - tests/
external/ , Python, 375 linesmrvi/ test_mrvi_model.py - tests/
external/ , Python, 85 linespoissonvi/ test_poissonvi.py - tests/
external/ , Python, 227 linesresolvi/ test_resolvi.py - tests/
external/ , Python, 27 linesscar/ test_scar.py - tests/
external/ , Python, 40 linesscbasset/ test_scbasset.py - tests/
external/ , Python, 478 linesscviva/ test_scviva.py - tests/
external/ , Python, 113 linessolo/ test_solo.py - tests/
external/ , Python, 42 linesstereoscope/ test_stereoscope.py - tests/
external/ , Python, 343 linessysvi/ test_sysvi.py - tests/
external/ , Python, 95 linestangram/ test_tangram.py - tests/
external/ , Python, 517 linestotalanvi/ test_totalanvi.py - tests/
external/ , Python, 65 linesvelovi/ test_velovi.py - tests/
external/ , Python, 467 linesvivs/ test_vivs.py - tests/
hub/ , Python, 196 linestest_hub_metadata.py - tests/
hub/ , Python, 399 linestest_hub_model.py - tests/
hub/ , Python, 49 linestest_url.py - tests/
model/ , Python, 90 linesbase/ test_base_model.py - tests/
model/ , Python, 33 linesbase/ test_rnamixin.py - tests/
model/ , Python, 124 linestest_amortizedlda.py - tests/
model/ , Python, 167 linestest_autozi.py - tests/
model/ , Python, 81 linestest_condscvi.py - tests/
model/ , Python, 194 linestest_destvi.py - tests/
model/ , Python, 163 linestest_differential.py - tests/
model/ , Python, 174 linestest_differential_abunda nce.py - tests/
model/ , Python, 38 linestest_docstrings.py - tests/
model/ , Python, 157 linestest_linear_scvi.py - tests/
model/ , Python, 224 linestest_mlxscvi.py - tests/
model/ , Python, 546 linestest_models_with_minifie d_data.py - tests/
model/ , Python, 512 linestest_models_with_mudata_ minified_data.py - tests/
model/ , Python, 338 linestest_multigpu.py - tests/
model/ , Python, 609 linestest_multivi.py - tests/
model/ , Python, 238 linestest_peakvi.py - tests/
model/ , Python, 562 linestest_pyro.py - tests/
model/ , Python, 741 linestest_scanvi.py - tests/
model/ , Python, 1,716 linestest_scvi.py - tests/
model/ , Python, 831 linestest_totalvi.py - tests/
model/ , Python, 33 linestest_utils.py - tests/
module/ , Python, 41 linestest_scanvae.py - tests/
module/ , Python, 71 linestest_vae.py - tests/
nn/ , Python, 76 linestest_embedding.py - tests/
nn/ , Python, 245 linestest_fclayers.py - tests/
test_warnings.py , Python, 46 lines - tests/
train/ , Python, 1 line__init__.py - tests/
train/ , Python, 267 linestest_callbacks.py - tests/
train/ , Python, 44 linestest_config.py - tests/
train/ , Python, 212 linestest_trainingplans.py - LICENSE, License, 30 lines
- README.md, Text, 147 lines
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 343 scripts, each with its path and the digest of its content;
- 1 match between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
Datasets cited
- zenodo:15320250, at Zenodo; found in “Data and materials availability”
Data and materials availability
The integrated C. robusta larva-stage single-cell RNA sequencing dataset and the complete differential expression data for all clusters shown in Figure 4C are deposited at Zenodo (https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 5 keywords, 7 MeSH terms, 2 funders, 104 references.
Cite
This paper
Kourakis, M. J., Miao, Y., Newman-Smith, E. D., Ryan, K., & Smith, W. C. (2026). Similarities between &
BibTeX
@article{kourakis2026sim
author = {Kourakis, Matthew J. and Miao, Yishen and Newman-Smith, Erin D. and Ryan, Kerrianne and Smith, William C.},
title = {{Similarities between \&
journal = {eNeuro},
year = {2026},
month = aug,
volume = {13},
number = {8},
pages = {ENEURO.0362--25.2026},
publisher = {Society for Neuroscience},
issn = {2373-2822},
doi = {10.1523/
url = {https://
pmid = {42637546},
pmcid = {PMC13509005}
}
RIS
TY - JOUR
AU - Kourakis, Matthew J.
AU - Miao, Yishen
AU - Newman-Smith, Erin D.
AU - Ryan, Kerrianne
AU - Smith, William C.
TI - Similarities between &
T2 - eNeuro
J2 - eNeuro
PY - 2026
DA - 2026/
VL - 13
IS - 8
SP - ENEURO.0362
EP - 25.2026
SN - 2373-2822
PB - Society for Neuroscience
DO - 10.1523/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1523/
"type": "article-journal",
"title": "Similarities between &
"container-title": "eNeuro",
"author": [
{
"family": "Kourakis",
"given": "Matthew J."
},
{
"family": "Miao",
"given": "Yishen"
},
{
"family": "Newman-Smith",
"given": "Erin D."
},
{
"family": "Ryan",
"given": "Kerrianne"
},
{
"family": "Smith",
"given": "William C."
}
],
"container-title-short":
"volume": "13",
"issue": "8",
"page": "ENEURO.0362-25.2026",
"DOI": "10.1523/
"PMID": "42637546",
"PMCID": "PMC13509005",
"ISSN": "2373-2822",
"publisher": "Society for Neuroscience",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
24
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1093/nar/gkag706 [code]
- scDifformer: diffusion-based post-training for virtual cell modeling across large-scale single-cell data.Journal: Nucleic acids researchIn common: PyTorch Geometric, anndata, Numba, 10 other tools, 1 reference
- [2] doi:10.1038/s41592-026-03057-2 [code]
- CREsted: modeling genomic and synthetic cell-type-specific enhancers across tissues and species.Journal: Nature methodsIn common: Biopython, anndata, Numba, 10 other tools, 1 reference
- [3] doi:10.1016/j.isci.2026.116055 [code]
- Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.Journal: iScienceIn common: PyTorch Geometric, anndata, Numba, 10 other tools, 1 reference
- [4] doi:10.1093/bioinformatics/btag652 [code]
- mmVelo: a deep generative model for estimating cell state-dependent dynamics across multiple modalities.Journal: Bioinformatics (Oxford, England)In common: Pyro, PyTorch Lightning, anndata, 9 other tools, 1 reference
- [5] doi:10.1038/s41586-026-10629-x [code]
- Whole-genome duplication shaped cell-type evolution in the vertebrate brain.Journal: NatureIn common: anndata, Numba, Scanpy, 7 other tools, 3 references
- [6] doi:10.1002/advs.77003 [code]
- SemanticST: A Scalable Multi-Contextual Graph Learning Framework for Uncovering Spatial Niches and Robust Multi-Sample Integration in Spatial Transcriptomics.Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)In common: PyTorch Geometric, anndata, Numba, 9 other tools, 1 reference
- [7] doi:10.1093/bioinformatics/btag540 [code]
- Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.Journal: Bioinformatics (Oxford, England)In common: PyTorch Geometric, anndata, Numba, 9 other tools, 1 reference
- [8] doi:10.1038/s42003-026-10957-8 [code]
- Brain defence by the extracellular matrix protein Cochlin.Journal: Communications biologyIn common: PyTorch Lightning, Biopython, SHAP, 9 other tools
- [9] doi:10.64898/2026.03.30.714220 [code]
- An integrated single cell and spatial omics atlas of human prenatal developmentJournal: bioRxiv (preprint)In common: Biopython, anndata, Scanpy, 9 other tools, 1 reference
- [10] doi:10.1038/s41467-026-74694-6 [code]
- Semi-supervised Omics Factor Analysis (SOFA) disentangles known and latent sources of variation in multi-omic data.Journal: Nature communicationsIn common: Pyro, anndata, Scanpy, 8 other tools, 1 reference
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 343 scripts, and 1 match between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:d601f3a71ed0da47…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
