From 1841abb648de9d1805dbcf0485615751a6d12bc7 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Sat, 21 Dec 2024 20:38:08 -0800 Subject: [PATCH 01/24] Add muANVI --- src/scvi/_constants.py | 1 + src/scvi/external/__init__.py | 5 + src/scvi/external/muanvi/__init__.py | 4 + src/scvi/external/muanvi/_base_components.py | 238 ++++++ src/scvi/external/muanvi/_model.py | 730 +++++++++++++++++++ src/scvi/external/muanvi/_module.py | 456 ++++++++++++ src/scvi/external/muanvi/_utils.py | 179 +++++ tests/external/muanvi/test_muanvi.py | 145 ++++ 8 files changed, 1758 insertions(+) create mode 100644 src/scvi/external/muanvi/__init__.py create mode 100644 src/scvi/external/muanvi/_base_components.py create mode 100644 src/scvi/external/muanvi/_model.py create mode 100644 src/scvi/external/muanvi/_module.py create mode 100644 src/scvi/external/muanvi/_utils.py create mode 100644 tests/external/muanvi/test_muanvi.py diff --git a/src/scvi/_constants.py b/src/scvi/_constants.py index ec6e4e914d..3dc4ea833b 100644 --- a/src/scvi/_constants.py +++ b/src/scvi/_constants.py @@ -5,6 +5,7 @@ class _REGISTRY_KEYS_NT(NamedTuple): X_KEY: str = "X" ATAC_X_KEY: str = "atac" BATCH_KEY: str = "batch" + SITE_KEY: str = "site" SAMPLE_KEY: str = "sample" LABELS_KEY: str = "labels" PROTEIN_EXP_KEY: str = "proteins" diff --git a/src/scvi/external/__init__.py b/src/scvi/external/__init__.py index 4621b23dad..4f7d67018c 100644 --- a/src/scvi/external/__init__.py +++ b/src/scvi/external/__init__.py @@ -11,7 +11,11 @@ from .gimvi import GIMVI from .methylvi import METHYLANVI, METHYLVI from .mrvi import MRVI +<<<<<<< HEAD from .mrvi_torch import TorchMRVI +======= +from .muanvi import MUANVI +>>>>>>> bb720b5f (Add muANVI) from .poissonvi import POISSONVI from .resolvi import RESOLVI from .scar import SCAR @@ -45,6 +49,7 @@ "SCVIVA", "CYTOVI", "DIAGVI", + "MUANVI", ] diff --git a/src/scvi/external/muanvi/__init__.py b/src/scvi/external/muanvi/__init__.py new file mode 100644 index 0000000000..fde12a9392 --- /dev/null +++ b/src/scvi/external/muanvi/__init__.py @@ -0,0 +1,4 @@ +from ._model import MUANVI +from ._module import MUANVAE + +__all__ = ["MUANVI", "MUANVAE"] diff --git a/src/scvi/external/muanvi/_base_components.py b/src/scvi/external/muanvi/_base_components.py new file mode 100644 index 0000000000..485040a764 --- /dev/null +++ b/src/scvi/external/muanvi/_base_components.py @@ -0,0 +1,238 @@ +import torch +from torch import nn as nn +from torch.nn import functional as F + +from scvi.module import Classifier +from scvi.nn import FCLayers + + +class Hierarchical_Classifier(nn.Module): + """ + Hierarchical Embedding Network + + Parameters (same as Classifier ) + ---------- + n_input + Number of input dimensions (dimensions of the latent space) + num_classes + number of labels in each label level in hierarchical list (ex : [2, 7]) + n_hidden + Number of hidden nodes in one layer + n_layers + Number of hidden layers per NN (per independent representation) + n_output + Number of dimensions of each independent representation + dropout_rate + dropout_rate for nodes + use_batch_norm + Whether to use batch norm in layers + use_layer_norm + Whether to use layer norm in layers + concatenation + Whether to concatenate or not the independent representations between layers + By default no concatenation. + """ + + def __init__( + self, + n_input: int, + num_classes: list, + n_hidden: int = 128, + dropout_rate: float = 0.1, + activation_fn: nn.Module = nn.ReLU, + n_layers: int = 3, + use_batch_norm: bool = False, + use_layer_norm: bool = True, + concatenation: bool = False, + logits: bool = True, # noqa, not used, we always return logits + ): + super().__init__() + self.n_input = n_input + self.n_hidden = n_hidden + self.concatenation = concatenation + + # independant representation level 1 of root level + layers = [ + FCLayers( + n_in=n_input, + n_out=n_hidden, + n_layers=n_layers, + n_hidden=n_hidden, + dropout_rate=dropout_rate, + use_batch_norm=use_batch_norm, + use_layer_norm=use_layer_norm, + activation_fn=activation_fn, + ) + for _ in range(len(num_classes)) + ] + # neural networks to obtain independant representations of dim n_output : + self.lvls = nn.ModuleList([nn.Sequential(layer) for layer in layers]) + + self.logits = nn.ModuleList([nn.Linear(n_hidden, num_class) for num_class in num_classes]) + self.softmax = nn.Softmax(dim=-1) + + def forward(self, x): + lvl_independents = [lvl(x) for lvl in self.lvls] + + logits_level = [ + logit(lvl_independent) + for logit, lvl_independent in zip(self.logits, lvl_independents, strict=False) + ] + probs_level = [self.softmax(logit_level) for logit_level in logits_level] + return probs_level, logits_level + + +class HierarchicalLossNetwork(Hierarchical_Classifier): + """ + Parameters (same as Classifier) + + ---------- + n_input + Number of input dimensions (dimensions of the latent space) + num_classes + number of labels in each class in hierarchical list (ex : [2, 7]) + n_hidden + Number of hidden nodes in one layer + n_layers + Number of hidden layers per NN (per independant representation) + n_output + Number of dimensions of each independant representation + dropout_rate + dropout_rate for nodes + use_batch_norm + Whether to use batch norm in layers + use_layer_norm + Whether to use layer norm in layers + activation_fn + Valid activation function from torch.nn + """ + + def __init__( + self, + n_input: int, + num_classes: list, + n_hidden: int = 128, + n_layers: int = 1, + dropout_rate: float = 0.1, + activation_fn: nn.Module = nn.ReLU, + use_batch_norm: bool = True, + use_layer_norm: bool = False, + **cls_parameters, + ): + # initialize the Classifier + super().__init__( + n_input=n_input, + n_hidden=n_hidden, + dropout_rate=dropout_rate, + activation_fn=activation_fn, + n_layers=n_layers, + use_batch_norm=use_batch_norm, + use_layer_norm=use_layer_norm, + num_classes=num_classes, + **cls_parameters, + ) + + self.total_level = len(num_classes) + + def calculate_lloss(self, predictions, true_labels, device, weights=None): + """ + Calculates the layer loss across all levels (multiple Cross-Entropy) + + Parameters + ---------- + predictions + Predictions of the model + true labels + Ground truth + weights + If we want to compute weighted cross entropy + """ + lloss = 0 + for l in range(self.total_level): + lloss += nn.CrossEntropyLoss()(predictions[l], true_labels[l]) + return lloss + + +class MultiBatchClassifier(nn.Module): + """ + Dictionary of batch-specific fully-connected NN classifiers. + + Parameters : parameters of the batch-specific classifiers + ---------- + n_input + Number of input dimensions + n_hidden + Number of nodes in hidden layer(s). If `0`, the classifier only consists of a + single linear layer. + n_labels + Numput of outputs dimensions + n_layers + Number of hidden layers. If `0`, the classifier only consists of a single + linear layer. + dropout_rate + dropout_rate for nodes + logits + Return logits or not + use_batch_norm + Whether to use batch norm in layers + use_layer_norm + Whether to use layer norm in layers + activation_fn + Valid activation function from torch.nn + n_sites + Number of different batch-specific classifiers to create + **kwargs + Keyword arguments passed into :class:`~scvi.nn.FCLayers`. + """ + + def __init__( + self, + n_input: int, + n_sites: int, + n_hidden: int = 128, + n_labels: int = 5, + n_layers: int = 1, + dropout_rate: float = 0.1, + use_batch_norm: bool = True, + use_layer_norm: bool = False, + activation_fn: nn.Module = nn.ReLU, + **kwargs, + ): + super().__init__() + self.n_sites = n_sites + self.classifier_dict = nn.ModuleDict( + { + str(i): Classifier( + n_input=n_input, + n_hidden=n_hidden, + n_labels=n_labels, + n_layers=n_layers, + dropout_rate=dropout_rate, + use_batch_norm=use_batch_norm, + use_layer_norm=use_layer_norm, + activation_fn=activation_fn, + **kwargs, + ) + for i in range(n_sites) + } + ) + + def forward(self, x, site_index): + """Forward computation for one mini batch of observations.""" + logits_list = [] + indices_list = [] + unique_sites = torch.unique(site_index) + + for site in unique_sites: + site_indices = torch.nonzero(site_index == site, as_tuple=True)[0].to(x.device) + indices_list.append(site_indices) + x_site = x[site_indices] + logits_list.append(self.classifier_dict[str(int(site.item()))](x_site)) + all_logits = torch.cat(logits_list, dim=0).to(x.device) + all_indices = torch.cat(indices_list, dim=0) + + # Sort the indices to get the original minibatch order + sorted_indices = torch.argsort(all_indices).to(x.device) + output_logits = all_logits[sorted_indices] + + return F.softmax(output_logits, dim=-1), output_logits \ No newline at end of file diff --git a/src/scvi/external/muanvi/_model.py b/src/scvi/external/muanvi/_model.py new file mode 100644 index 0000000000..e2b6e87531 --- /dev/null +++ b/src/scvi/external/muanvi/_model.py @@ -0,0 +1,730 @@ +import logging +import warnings +from collections.abc import Sequence +from copy import deepcopy +from typing import Literal + +import numpy as np +import pandas as pd +import torch +from anndata import AnnData + +from scvi import REGISTRY_KEYS +from scvi.data import AnnDataManager +from scvi.data._constants import _SETUP_ARGS_KEY +from scvi.data.fields import ( + CategoricalJointObsField, + CategoricalObsField, + LabelsWithUnlabeledObsField, + LayerField, + NumericalJointObsField, + NumericalObsField, +) +from scvi.dataloaders import SemiSupervisedDataSplitter +from scvi.model import SCVI +from scvi.model._utils import get_max_epochs_heuristic, parse_device_args +from scvi.model.base import ArchesMixin, BaseModelClass, RNASeqMixin, VAEMixin +from scvi.model.base._archesmixin import _get_loaded_data, _set_params_online_update +from scvi.model.base._save_load import ( + _initialize_model, + _validate_var_names, +) +from scvi.train import SemiSupervisedTrainingPlan, TrainRunner +from scvi.train._callbacks import SubSampleLabels +from scvi.utils import setup_anndata_dsp +from scvi.utils._docstrings import devices_dsp + +from ._module import MUANVAE +from ._utils import LabelsWithUnlabeledJointObsField, _get_site_code_from_category + +logger = logging.getLogger(__name__) + + +class MUANVI(RNASeqMixin, VAEMixin, ArchesMixin, BaseModelClass): + """ + Hierarchical multi-annotator Variational Inference [Xu21]_. + + Inspired from M1 + M2 model, as described in (https://arxiv.org/pdf/1406.5298.pdf). + + Parameters + ---------- + adata + AnnData object that has been registered via :meth:`~scvi.model.MUANVI.setup_anndata`. + n_hidden + Number of nodes per hidden layer. + n_latent + Dimensionality of the latent space. + n_layers + Number of hidden layers used for encoder and decoder 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 + gene_likelihood + One of: + * ``'nb'`` - Negative binomial distribution + * ``'zinb'`` - Zero-inflated negative binomial distribution + * ``'poisson'`` - Poisson distribution + update_yprior + Whether to perform the hierarchical update of the y prior parameter in the loss + batches_to_harmonize + List of two indices of the two batches to label-harmonize. Has to be defined if the dataset has more than 2 batches. + **model_kwargs + Keyword args for :class:`~scvi.module.MUANVAE` + + Examples + -------- + >>> adata = anndata.read_h5ad(path_to_anndata) + >>> scvi.external.MUANVI.setup_anndata(adata, labels=["labels_0", "labels_1"], unknown_categories=["unknown, "unknown"]) + >>> model = scvi.external.MUANVI(adata) + >>> model.train() + >>> adata.obsm["X_muanvi"] = model.get_latent_representation() + >>> adata.obs["pred_label_coarse"] = model.predict()[0] + >>> adata.obs["pred_label_fine"] = model.predict()[1] + + """ + + _module_cls = MUANVAE + _training_plan_cls = SemiSupervisedTrainingPlan + + def __init__( + self, + adata: AnnData, + n_hidden: int = 128, + n_latent: int = 10, + n_layers: int = 1, + dropout_rate: float = 0.1, + dispersion: Literal["gene", "gene-batch", "gene-label", "gene-cell"] = "gene", + gene_likelihood: Literal["zinb", "nb", "poisson"] = "nb", + update_yprior: bool = True, + eps_yprior: float = 1e-4, + **model_kwargs, + ): + super().__init__(adata) + muanvae_model_kwargs = dict(model_kwargs) + self._set_indices_and_labels() + + n_batch = self.summary_stats.n_batch + n_fine_labels = self.summary_stats.n_labels - 1 + print(self.summary_stats) + n_site = self.summary_stats.n_site + + hierarchy_dict, num_classes, hierarchy_matrix = self.extract_hierarchy( + n_site=n_site, eps_yprior=eps_yprior + ) + n_cats_per_cov = ( + self.adata_manager.get_state_registry(REGISTRY_KEYS.CAT_COVS_KEY).n_cats_per_key + if REGISTRY_KEYS.CAT_COVS_KEY in self.adata_manager.data_registry + else None + ) + + use_size_factor_key = REGISTRY_KEYS.SIZE_FACTOR_KEY in self.adata_manager.data_registry + + self.module = self._module_cls( + n_input=self.summary_stats.n_vars, + n_batch=n_batch, + n_site=n_site, + n_fine_labels = n_fine_labels, + num_classes=num_classes, + n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), + n_cats_per_cov=n_cats_per_cov, + n_hidden=n_hidden, + n_latent=n_latent, + n_layers=n_layers, + dropout_rate=dropout_rate, + dispersion=dispersion, + gene_likelihood=gene_likelihood, + use_size_factor_key=use_size_factor_key, + hierarchy_dict=hierarchy_dict, + hierarchy_matrix=hierarchy_matrix, + update_yprior=update_yprior, + **muanvae_model_kwargs, + ) + + self.unsupervised_history_ = None + self.semisupervised_history_ = None + + self._model_summary_string = ( + f"muANVI Model with the following params: \nunlabeled_category: {self.unlabeled_category}, n_hidden: {n_hidden}, n_latent: {n_latent}" + f", n_layers: {n_layers}, dropout_rate: {dropout_rate}, dispersion: {dispersion}, gene_likelihood: {gene_likelihood}" + ) + self.init_params_ = self._get_init_params(locals()) + self.was_pretrained = False + self.n_fine_labels = n_fine_labels + + @classmethod + def from_scvi_model( + cls, + scvi_model: SCVI, + unlabeled_category: list[str], + fine_labels: str | None = None, + label_hierarchy: list[str] | None = None, + adata: AnnData | None = None, + **muanvi_kwargs, + ): + """ + Initialize scHANVI model with weights from pretrained :class:`~scvi.model.SCVI` model. + + Parameters + ---------- + scvi_model + Pretrained scvi model + fine_labels + key in `adata.obs` for label information. If this value is not None, the key will + overwrite the `labels_key` used to setup AnnData with scvi. + label_hierarchy + List of strings, with levels of the cell-type hierarchy. Full hierarchy is inferred by + concatenating label_hierarchy and fine_labels. + unlabeled_category + Value used for unlabeled cells in `labels_key`. + adata + AnnData object that has been registered via :meth:`~scvi.model.MUANVI.setup_anndata`. + muanvi_kwargs + kwargs for muANVI model + """ + scvi_model._check_if_trained(message="Passed in scvi model hasn't been trained yet.") + + muanvi_kwargs = dict(muanvi_kwargs) + init_params = scvi_model.init_params_ + non_kwargs = init_params["non_kwargs"] + kwargs = init_params["kwargs"] + kwargs = {k: v for (i, j) in kwargs.items() for (k, v) in j.items()} + for k, v in {**non_kwargs, **kwargs}.items(): + if k in muanvi_kwargs.keys(): + warnings.warn( + f"Ignoring param '{k}' as it was already passed in to " + + f"pretrained scvi model with value {v}.", + stacklevel=2, + ) + del muanvi_kwargs[k] + + if adata is None: + adata = scvi_model.adata + else: + # validate new anndata against old model + scvi_model._validate_anndata(adata) + + scvi_setup_args = deepcopy(scvi_model.adata_manager.registry[_SETUP_ARGS_KEY]) + scvi_labels_key = scvi_setup_args["labels_key"] + if fine_labels is None and scvi_labels_key is None: + raise ValueError( + "A `labels_key` list is necessary as the SCVI model was initialized without one." + ) + if fine_labels is not None: + scvi_setup_args.update({"fine_labels": fine_labels}) + scvi_setup_args.pop("labels_key", None) + else: + scvi_setup_args["fine_labels"] = scvi_setup_args.pop("labels_key") + + cls.setup_anndata( + adata, + unlabeled_category=unlabeled_category, + label_hierarchy=label_hierarchy, + **scvi_setup_args, + ) + muanvi_model = cls(adata, **non_kwargs, **kwargs, **muanvi_kwargs) + scvi_state_dict = scvi_model.module.state_dict() + muanvi_model.module.load_state_dict(scvi_state_dict, strict=False) + muanvi_model.was_pretrained = True + + return muanvi_model + + def _set_indices_and_labels(self): + """Set indices for labeled and unlabeled cells.""" + labels_state_registry = self.adata_manager.get_state_registry("label_hierarchy") + self.original_label_keys = labels_state_registry.field_keys + self.unlabeled_category = labels_state_registry.unlabeled_category + + # Dataframe of 2 columns for the 2 layers of labels + labels = {field: self.adata.obs[field] for field in self.original_label_keys} + self.labels = pd.DataFrame(labels) + self._label_mapping = labels_state_registry.mappings.to_dict() + # a cell is unlabeled if it is not labeled at the finest state + labeled_indices_list = [ + set(np.where(self.labels.iloc[:, idx] != self.unlabeled_category)[0]) + for idx, _ in enumerate(self.original_label_keys) + ] + self._labeled_indices = list(set.intersection(*labeled_indices_list)) + self._unlabeled_indices = list( + set(np.arange(self.adata.n_obs)) - set(self._labeled_indices) + ) + self._code_to_label = [ + dict(enumerate(self._label_mapping[layer])) for layer in self._label_mapping + ] + + def extract_hierarchy(self, n_site, eps_yprior): + """ + Method to extract automatically the intrinsic hierarchy in the data. + + Parameters + ---------- + n_site + Number of different vocabulary sites in the data + eps_yprior + Epsilon value to add to the y prior parameter to avoid overconfidence + """ + labels_state_registry = self.adata_manager.get_state_registry("label_hierarchy") + label_keys = labels_state_registry.field_keys + + def nested_groupby(df, label_keys): + # Recursively adds dicts for each level of hierarchy. + if len(label_keys) == 1: + return df.groupby(label_keys[0]).apply(list).to_dict() + else: + return { + key: nested_groupby(sub_df, label_keys[1:]) + for key, sub_df in df.groupby(label_keys[0]) + } + + hierarchy_dict = nested_groupby(self.labels, label_keys) + num_classes = labels_state_registry.n_cats_per_key + hierarchy_matrix = [] + + for n_label in range(1, len(num_classes)): + # Site specific last layer. + if n_label == len(num_classes) - 1: + hierarchy_matrix_ = torch.zeros( + num_classes[n_label - 1], num_classes[n_label], n_site + ) + curr = pd.DataFrame( + 0, + index=np.arange(hierarchy_matrix_.shape[0]), + columns=np.arange(hierarchy_matrix_.shape[1]), + ) + for site in range(n_site): + adata = self.adata[self.adata.obs["_scvi_site"] == site] + + if adata.n_obs > 0: + curr_ = pd.crosstab( + adata.obsm["_scvi_label_hierarchy"][ + label_keys[n_label - 1] + ], + adata.obsm["_scvi_label_hierarchy"][label_keys[n_label]], + ) + curr_[curr_ > 0] = 1 + curr_ = curr_.loc[curr.index, curr.columns] + curr = curr_.div(curr_.sum(axis=1), axis=0) + hierarchy_matrix_[:, :, site] = torch.tensor(curr.fillna(0).values) + else: + curr_ = pd.crosstab( + self.adata.obsm["_scvi_label_hierarchy"][label_keys[n_label - 1]], + self.adata.obsm["_scvi_label_hierarchy"][label_keys[n_label]], + ) + curr_[curr_ > 0] = 1 + curr_ = curr_.loc[curr.index, curr.columns] + curr = curr_[curr_.index, curr_.columns].div(curr_.sum(axis=1), axis=0).values + hierarchy_matrix_ = torch.tensor(curr.fillna(0).values) + + hierarchy_matrix.append(hierarchy_matrix_) + + return (hierarchy_dict, num_classes, hierarchy_matrix) + + def predict( + self, + adata: AnnData | None = None, + indices: Sequence[int] | None = None, + soft: bool = False, + batch_size: int | None = None, + level: int = -1, + sites_to_predict: int | str | str | None = None, + ) -> np.ndarray | pd.DataFrame: + """ + Return cell label predictions. + + Parameters + ---------- + adata + AnnData object that has been registered via :meth:`~scvi.model.SCANVI.setup_anndata`. + indices + indices for which to return probabilities. + soft + If True, returns per class probabilities + batch_size + Minibatch size for data loading into model. Defaults to `scvi.settings.batch_size`. + site_to_predict + For cross prediction purposes : sites to use when cross-predicting labels. + If None, normal prediction occurs and each cell is labeled accordingly to + its own site-specific classifier. Can be list of strings to use multiple sites, + or a string to use a single site. + level + Level of the hierarchy to predict. If -1, predicts at the finest level. + """ + adata = self._validate_anndata(adata) + + if indices is None: + indices = np.arange(adata.n_obs) + + scdl = self._make_data_loader( + adata=adata, + indices=indices, + batch_size=batch_size, + ) + # total depth of the hierarchy + total_level = len(self.module.num_classes) + class_labels = self.adata_manager.get_state_registry("label_hierarchy").field_keys + + sites_to_predict_, site_mappings_ = _get_site_code_from_category( + self.get_anndata_manager(adata, required=True), sites_to_predict + ) + pred = {str(site_mappings_[i]) + "_fine": [] for i in sites_to_predict_ if i is not None} + if None in sites_to_predict_: + for i in class_labels: + pred[i] = [] + + for _, tensors in enumerate(scdl): + for site_to_predict_ in sites_to_predict_: + probs_, _ = self.module.classification(tensors, site_to_predict=site_to_predict_) + + if site_to_predict_ is not None: + if not soft: + pred_ = probs_.argmax(dim=1) + else: + pred_ = probs_ + pred[site_mappings_[site_to_predict_] + "_fine"].append(pred_.detach().cpu()) + + else: + # pred is a tuple of probabilities and logits for each layer + probs = probs_ + + for i in range(total_level): + if not soft: + probs[i] = probs[i].argmax(dim=1) # select only the label + else: + pred[class_labels[i]].append(probs[i].detach().cpu()) + + for key in pred.keys(): + pred[key] = torch.cat(pred[key]).numpy() + if not soft: + if key not in class_labels: + pred[key] = [self._code_to_label[-1][ct] for ct in pred[key]] + else: + pred[key] = [ + self._code_to_label[class_labels.index(key)][ct] for ct in pred[key] + ] + + if not soft: + pred = pd.DataFrame.from_dict(pred) + pred.index = adata.obs_names[indices] + return pred + else: + for key in pred.keys(): + if key not in class_labels: + columns = list(self._code_to_label[-1].values())[:-1] + else: + columns = list(self._code_to_label[class_labels.index(key)].values())[:-1] + + pred[key] = pd.DataFrame( + pred[key], + columns=columns, + index=adata.obs_names[indices], + ) + return pred + + @classmethod + @devices_dsp.dedent + def load_query_data( + cls, + adata: AnnData, + reference_model: str | BaseModelClass, + inplace_subset_query_vars: bool = False, + accelerator: str = "auto", + device: int | str = "auto", + unfrozen: bool = False, + freeze_dropout: bool = False, + freeze_expression: bool = True, + freeze_decoder_first_layer: bool = True, + freeze_batchnorm_encoder: bool = True, + freeze_batchnorm_decoder: bool = False, + freeze_classifier: bool = True, + ): + """Online update of a reference model with scArches algorithm :cite:p:`Lotfollahi21`. + + Parameters + ---------- + adata + AnnData organized in the same way as data used to train model. + It is not necessary to run setup_anndata, + as AnnData is validated against the ``registry``. + reference_model + Either an already instantiated model of the same class, or a path to + saved outputs for reference model. + inplace_subset_query_vars + Whether to subset and rearrange query vars inplace based on vars used to + train reference model. + %(param_accelerator)s + %(param_device)s + unfrozen + Override all other freeze options for a fully unfrozen model + freeze_dropout + Whether to freeze dropout during training + freeze_expression + Freeze neurons corersponding to expression in first layer + freeze_decoder_first_layer + Freeze neurons corresponding to first layer in decoder + freeze_batchnorm_encoder + Whether to freeze batchnorm weight and bias during training for encoder + freeze_batchnorm_decoder + Whether to freeze batchnorm weight and bias during training for decoder + freeze_classifier + Whether to freeze classifier completely. Only applies to `SCANVI`. + """ + _, _, device = parse_device_args( + accelerator=accelerator, + devices=device, + return_device="torch", + validate_single_device=True, + ) + + attr_dict, var_names, load_state_dict = _get_loaded_data(reference_model, device=device) + + if inplace_subset_query_vars: + logger.debug("Subsetting query vars to reference vars.") + adata._inplace_subset_var(var_names) + _validate_var_names(adata, var_names) + + registry = attr_dict.pop("registry_") + if _SETUP_ARGS_KEY not in registry: + raise ValueError( + "Saved model does not contain original setup inputs. " + "Cannot load the original setup." + ) + + cls.setup_anndata( + adata, + source_registry=registry, + extend_categories=True, + allow_missing_labels=True, + **registry[_SETUP_ARGS_KEY], + ) + + model = _initialize_model(cls, adata, attr_dict) + adata_manager = model.get_anndata_manager(adata, required=True) + + if REGISTRY_KEYS.CAT_COVS_KEY in adata_manager.data_registry: + raise NotImplementedError( + "scArches currently does not support models with extra categorical covariates." + ) + + model.to_device(device) + + # model tweaking + new_state_dict = model.module.state_dict() + additional_parameters = set() + for key, new_ten in new_state_dict.items(): + load_ten = load_state_dict.get(key, None) + if load_ten is None: + # Picks up that additional site classifier was added, makes it trainable by default. + if "y_prior_fine" not in key: + additional_parameters.add(key) # TODO check this. + load_state_dict[key] = new_ten + continue + if new_ten.size() == load_ten.size(): + continue + # new categoricals changed size + else: + if new_ten.size()[0] != load_ten.size()[0]: + new_ten = new_ten.to(load_ten.device) + dim_diff = new_ten.size()[0] - load_ten.size()[0] + fixed_ten = torch.cat([load_ten, new_ten[-dim_diff:, ...]], dim=0) + load_state_dict[key] = fixed_ten + else: + new_ten = new_ten.to(load_ten.device) + dim_diff = new_ten.size()[-1] - load_ten.size()[-1] + fixed_ten = torch.cat([load_ten, new_ten[..., -dim_diff:]], dim=-1) + load_state_dict[key] = fixed_ten + + model.module.load_state_dict(load_state_dict) + model.module.eval() + + _set_params_online_update( + model.module, + unfrozen=unfrozen, + freeze_decoder_first_layer=freeze_decoder_first_layer, + freeze_batchnorm_encoder=freeze_batchnorm_encoder, + freeze_batchnorm_decoder=freeze_batchnorm_decoder, + freeze_dropout=freeze_dropout, + freeze_expression=freeze_expression, + freeze_classifier=freeze_classifier, + parameters_yes_grad=additional_parameters, + ) + model.is_trained_ = False + + return model + + def train( + self, + max_epochs: int | None = None, + n_samples_per_label: float | None = None, + check_val_every_n_epoch: int | None = None, + train_size: float = 0.9, + validation_size: float | None = None, + shuffle_set_split: bool = True, + batch_size: int = 128, + accelerator: str = "auto", + devices: int | list[int] | str = "auto", + datasplitter_kwargs: dict | None = None, + plan_kwargs: dict | None = None, + **trainer_kwargs, + ): + """Train the model. + + Parameters + ---------- + max_epochs + Number of passes through the dataset for semisupervised training. + n_samples_per_label + Number of subsamples for each label class to sample per epoch. By default, there + is no label subsampling. + check_val_every_n_epoch + Frequency with which metrics are computed on the data for validation set for both + the unsupervised and semisupervised trainers. If you'd like a different frequency for + the semisupervised trainer, set check_val_every_n_epoch in semisupervised_train_kwargs. + train_size + Size of training set in the range [0.0, 1.0]. + validation_size + Size of the test set. If `None`, defaults to 1 - `train_size`. If + `train_size + validation_size < 1`, the remaining cells belong to a test set. + shuffle_set_split + Whether to shuffle indices before splitting. If `False`, the val, train, and test set + are split in the sequential order of the data according to `validation_size` and + `train_size` percentages. + batch_size + Minibatch size to use during training. + %(param_accelerator)s + %(param_devices)s + datasplitter_kwargs + Additional keyword arguments passed into + :class:`~scvi.dataloaders.SemiSupervisedDataSplitter`. + plan_kwargs + Keyword args for :class:`~scvi.train.SemiSupervisedTrainingPlan`. Keyword arguments + passed to `train()` will overwrite values present in `plan_kwargs`, when appropriate. + **trainer_kwargs + Other keyword args for :class:`~scvi.train.Trainer`. + """ + if max_epochs is None: + max_epochs = get_max_epochs_heuristic(self.adata.n_obs) + + if self.was_pretrained: + max_epochs = int(np.min([10, np.max([2, round(max_epochs / 3.0)])])) + + plan_kwargs = {} if plan_kwargs is None else plan_kwargs + datasplitter_kwargs = datasplitter_kwargs or {} + + # if we have labeled cells, we want to subsample labels each epoch + sampler_callback = [SubSampleLabels()] if len(self._labeled_indices) != 0 else [] + + data_splitter = SemiSupervisedDataSplitter( + adata_manager=self.adata_manager, + train_size=train_size, + validation_size=validation_size, + shuffle_set_split=shuffle_set_split, + n_samples_per_label=n_samples_per_label, + batch_size=batch_size, + **datasplitter_kwargs, + ) + + warmup_epochs = plan_kwargs.pop("warmup_epochs", None) + + if warmup_epochs is not None and warmup_epochs > 0: + logger.info(f"Pretraining for {max_epochs} epochs.") + + plan_kwargs_pre = plan_kwargs.copy() + plan_kwargs_pre["warmup_model"] = True + plan_kwargs_pre["n_epochs_kl_warmup"] = warmup_epochs + + training_plan = self._training_plan_cls( + self.module, n_classes=self.n_fine_labels, **plan_kwargs_pre + ) # n_classes set at the finest level to track accuracy at that level. + runner_pre = TrainRunner( + self, + training_plan=training_plan, + data_splitter=data_splitter, + max_epochs=warmup_epochs, + accelerator=accelerator, + devices=devices, + check_val_every_n_epoch=check_val_every_n_epoch, + **trainer_kwargs, + ) + runner_pre() + self.was_pretrained = True + + logger.info(f"Training for {max_epochs} epochs.") + + if "callbacks" in trainer_kwargs.keys(): + trainer_kwargs["callbacks"] + [sampler_callback] + else: + trainer_kwargs["callbacks"] = sampler_callback + training_plan = self._training_plan_cls( + self.module, n_classes=self.n_fine_labels, **plan_kwargs + ) + + runner = TrainRunner( + self, + training_plan=training_plan, + data_splitter=data_splitter, + max_epochs=max_epochs, + accelerator=accelerator, + devices=devices, + check_val_every_n_epoch=check_val_every_n_epoch, + **trainer_kwargs, + ) + + return runner() + + @classmethod + @setup_anndata_dsp.dedent + def setup_anndata( + cls, + adata: AnnData, + fine_labels_key: str, + unlabeled_category: list[str | int | float], + layer: str | None = None, + site_key: str | None = None, + batch_key: str | None = None, + size_factor_key: str | None = None, + categorical_covariate_keys: list[str] | None = None, + continuous_covariate_keys: list[str] | None = None, + label_hierarchy: list[str] | None = None, + **kwargs, + ): + """ + %(summary)s. + + Parameters + ---------- + %(param_layer)s + %(param_batch_key)s + %(param_site_key)s + fine_labels_key + key in `adata.obs` for fine label information. Categories will automatically be + converted into integer categories and saved to `adata.obs['_scvi_labels']`. + If `None`, assigns the same label to all the data. This information can be + site-specific. In this case we expect all the labels in a single obs column. + This is analogous to the label key in scANVI. + %(param_size_factor_key)s + %(param_cat_cov_keys)s + %(param_cont_cov_keys)s + label_hierarchy + List of strings, with levels of the cell-type hierarchy. + The first list is the root level, the last list is the second finest level. + The full hierarchy is inferred by concatenating label_hierarchy and fine_labels. + """ + setup_method_args = cls._get_setup_method_args(**locals()) + anndata_fields = [ + LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), + CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), + CategoricalObsField(REGISTRY_KEYS.SITE_KEY, site_key), + LabelsWithUnlabeledObsField(REGISTRY_KEYS.LABELS_KEY, fine_labels_key, unlabeled_category), + NumericalObsField(REGISTRY_KEYS.SIZE_FACTOR_KEY, size_factor_key, required=False), + CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), + NumericalJointObsField(REGISTRY_KEYS.CONT_COVS_KEY, continuous_covariate_keys), + LabelsWithUnlabeledJointObsField( + "label_hierarchy", label_hierarchy+[fine_labels_key], unlabeled_category), + ] + adata_manager = AnnDataManager(fields=anndata_fields, setup_method_args=setup_method_args) + adata_manager.register_fields(adata, **kwargs) + cls.register_manager(adata_manager) diff --git a/src/scvi/external/muanvi/_module.py b/src/scvi/external/muanvi/_module.py new file mode 100644 index 0000000000..3fe1b5a30e --- /dev/null +++ b/src/scvi/external/muanvi/_module.py @@ -0,0 +1,456 @@ +from typing import Literal + +import torch +from torch.distributions import Categorical, Independent, MixtureSameFamily, Normal +from torch.distributions import kl_divergence as kl +from torch.nn import functional as F + +from scvi import REGISTRY_KEYS +from scvi.module import SCANVAE +from scvi.module._utils import broadcast_labels +from scvi.module.base import LossOutput, auto_move_data + +from ._base_components import HierarchicalLossNetwork, MultiBatchClassifier + + +class MUANVAE(SCANVAE): + """ + Single-cell multiple-annotation using variational inference. + + This is a re-implementation of a hierarchical cell-type annotation model + inspired from scANVI model described in [Xu21]_,. + + Parameters + ---------- + n_input + Number of input genes + n_batch + Number of batches + n_fine_labels + Number of fine labels + num_classes + Number of labels per class organized in a hierarchical list + n_hidden + Number of nodes per hidden layer + n_latent + Dimensionality of the latent space + n_layers + Number of hidden layers used for encoder and decoder NNs + n_continuous_cov + Number of continuous covariates + n_cats_per_cov + Number of categories for each extra categorical covariate + 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 + conditioning_class + index of the class conditioning the second latent space (0 being the coarse class, 1 being the fine class). Default : coarse labels. + hierarchy_matrix + Matrix representing the hierarchy, computed by scATVI. If None, no hierarchical y prior update in the loss. + use_batch_norm + Whether to use batch norm in layers + use_layer_norm + Whether to use layer norm in layers + prior_z1 + Whether to use MoG or simple Gaussian for prior of z1 + **vae_kwargs + Keyword args for :class:`~scvi.module.VAE` + """ + + def __init__( + self, + n_input: int, + num_classes: list, + hierarchy_dict: dict, + n_batch: int = 0, + n_site: int = 0, + n_fine_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: str = "gene", + log_variational: bool = True, + gene_likelihood: str = "nb", + classifier_parameters: dict = dict(), + classifier_parameters_muanvae: dict = dict(), + use_batch_norm: Literal["encoder", "decoder", "none", "both"] = "none", + use_layer_norm: Literal["encoder", "decoder", "none", "both"] = "both", + conditioning_class: int = -1, + mog_class: int = 0, + hierarchy_matrix=None, + update_yprior=True, + prior_z1: str = "gaussian", + eps_yprior: float = 1e-6, + **scanvae_kwargs, + ): + self.conditioning_class = conditioning_class + self.mog_class = mog_class + self.site_specific_classifier = n_site > 1 + self.n_site = n_site + self.num_classes = num_classes + self.update_yprior = update_yprior + self.n_labels_conditioning = num_classes[self.conditioning_class] + self.n_fine_labels = n_fine_labels + self.hiearchy_dict = hierarchy_dict + + classifier_parameters = classifier_parameters or {} + + cls_parameters = { + "n_layers": n_layers, + "n_hidden": n_hidden, + "dropout_rate": dropout_rate, + "logits": True, + } + cls_parameters.update(classifier_parameters) + + super().__init__( + n_input, + n_batch=n_batch, + n_labels=self.n_fine_labels, + n_hidden=n_hidden, + n_latent=n_latent, + n_layers=n_layers, + n_continuous_cov=n_continuous_cov, + n_cats_per_cov=n_cats_per_cov, + dropout_rate=dropout_rate, + dispersion=dispersion, + log_variational=log_variational, + gene_likelihood=gene_likelihood, + classifier_parameters=classifier_parameters, + use_batch_norm=use_batch_norm, + use_layer_norm=use_layer_norm, + **scanvae_kwargs, + ) + cls_parameters.update(classifier_parameters_muanvae) + self.num_classes = num_classes + self.total_level = len(self.num_classes) # depth of hierarchy + self.prior_z1 = prior_z1 + + self.classifier = HierarchicalLossNetwork( + n_input=self.n_latent, + num_classes=self.num_classes[:-1], + **cls_parameters, + ) + # the site-specific classifiers must have layer norm (batch norm does not work if there is 1 single observation from a batch in a minibatch) + self.multi_classifier_fine = MultiBatchClassifier( + n_input=self.n_latent, + n_sites=n_site, + n_labels=n_fine_labels, + use_batch_norm=False, + use_layer_norm=True, + **cls_parameters, + ) + + # register y_prior on the fine labels + if not self.update_yprior: + hierarchy_matrix = [ + torch.tensor(hierarchy_matrix[i].sum(0) > 0, dtype=torch.float) + for i in range(len(num_classes) - 1) + ] + for i in range(0, self.total_level): + if i == self.total_level - 1: + self.y_prior_fine = torch.nn.ParameterList( + [ + torch.nn.Parameter( # + hierarchy_matrix[i - 1][:, :, site] + eps_yprior, + requires_grad=False, + ) + for site in range(n_site) + ] + ) + elif i == 0: + self.register_buffer( + f"y_prior_{i}", + torch.nn.Parameter( # + torch.full([num_classes[0]], 1 / num_classes[0], dtype=torch.float), + requires_grad=False, + ), + ) + else: + self.register_buffer( + f"y_prior_{i}", + torch.nn.Parameter( # + torch.tensor(hierarchy_matrix[i - 1] + eps_yprior, dtype=torch.float), + requires_grad=False, + ), + ) + if self.prior_z1 == "mog": + self.register_parameter( + "prior_z1_means", + torch.nn.Parameter(torch.randn([num_classes[self.mog_class], n_latent])), + ) + self.register_parameter( + "prior_z1_scales", + torch.nn.Parameter(torch.zeros([num_classes[self.mog_class], n_latent])), + ) + self.register_parameter( + "prior_z1_logits", torch.nn.Parameter(torch.ones([num_classes[self.mog_class]])) + ) + if self.prior_z1 == "mog_celltype": + self.register_parameter( + "prior_z1_means", + torch.nn.Parameter(torch.zeros([1, num_classes[self.mog_class], n_latent])), + ) + self.register_parameter( + "prior_z1_scales", + torch.nn.Parameter(torch.zeros([1, num_classes[self.mog_class], n_latent])), + ) + + @auto_move_data + def classify( + self, + x, + batch_index=None, + site_index=None, + cont_covs=None, + cat_covs=None, + site_to_predict: int | None = None, + precomputed_z: torch.Tensor | None = None, + ): + """ + Classify cells using the model. + + Parameters + ---------- + site_to_predict + For cross prediction purposes : index of the fine classifier to use when cross-predicting labels. + If None, normal prediction occurs, which is the case during training. + If not None, cross classification on only fine layer with this batch-specific classifier. + precomputed_z + Precomputed z1 latent space. If None, z1 is computed from x. + """ + if precomputed_z is not None: + z = precomputed_z + else: + if self.log_variational: # for numerical stability + x = torch.log(1 + 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 = () + qz, _ = self.z_encoder( + encoder_input, batch_index, *categorical_input + ) # q(z1|x) without the var qz_v + z = qz.rsample() + if site_to_predict is not None: + site_to_predict_index = torch.full((z.shape[0], 1), site_to_predict, dtype=torch.int32) + probs_fine, _ = self.multi_classifier_fine(z, site_to_predict_index) + return probs_fine, _ + probs_classifier, logits_classifier = self.classifier(z) + probs_fine, logits_fine = self.multi_classifier_fine(z, site_index) + probs_classifier += [probs_fine] + logits_classifier += [logits_fine] + + return probs_classifier, logits_classifier + + @auto_move_data + def classification( + self, + tensors, + return_classifier_loss=False, + site_to_predict=None, + precomputed_z=None, + ): + x = tensors[REGISTRY_KEYS.X_KEY] + y = tensors["label_hierarchy"] + batch_idx = tensors[REGISTRY_KEYS.BATCH_KEY] + site_idx = tensors[REGISTRY_KEYS.SITE_KEY] + cont_covs = tensors.get(REGISTRY_KEYS.CONT_COVS_KEY, None) + cat_covs = tensors.get(REGISTRY_KEYS.CAT_COVS_KEY, None) + + probs, logits = self.classify( + x, + batch_index=batch_idx, + site_index=site_idx, + cat_covs=cat_covs, + cont_covs=cont_covs, + site_to_predict=site_to_predict, + precomputed_z=precomputed_z, + ) + if not return_classifier_loss: + return probs, logits + + classification_loss = 0 + for l in range(self.total_level): + labels_curr_level = y[:, l].view(-1).long() + classification_loss += F.cross_entropy( + logits[l], + labels_curr_level, + ignore_index=self.num_classes[l], + reduction="mean" + ) + true_labels_fine = torch.unsqueeze(labels_curr_level, 1) + return classification_loss, true_labels_fine, logits[-1] + + def loss( + self, + tensors, + inference_outputs, + generative_ouputs, + kl_weight=1, + labelled_tensors=None, + classification_ratio=None, + bg_classifier_ratio=0.0, + loss_z1_factor=0.1, + weighting_mog=1.0, + warmup_model=False, # if true, trains a model without cell-type classification first. + ): + """Compute the loss.""" + px = generative_ouputs["px"] + qz1 = inference_outputs["qz"] + z1 = inference_outputs["z"] + x = tensors[REGISTRY_KEYS.X_KEY] + y = tensors[REGISTRY_KEYS.LABELS_KEY] + site_index = tensors[REGISTRY_KEYS.SITE_KEY] + + is_labelled = False if y is None else True + + # Enumerate choices of label + ys, z1s = broadcast_labels(z1, n_broadcast=self.n_labels_conditioning) + qz2, z2 = self.encoder_z2_z1(z1s, ys) + pz1_m, pz1_v = self.decoder_z1_z2(z2, ys) + reconst_loss = -px.log_prob(x).sum(-1) + + # KL Divergence + mean = torch.zeros_like(qz2.loc) + scale = torch.ones_like(qz2.scale) + + kl_divergence_z2 = kl(qz2, Normal(mean, scale)).sum(dim=1) + loss_z1_unweight = -Normal(pz1_m, torch.sqrt(pz1_v)).log_prob(z1s).sum(dim=-1) + loss_z1_weight = qz1.log_prob(z1).sum(dim=-1) + + probs, logits = self.classification(tensors, precomputed_z=z1) + probs_conditioning, logits_conditioning = ( + probs[self.conditioning_class], + logits[self.conditioning_class], + ) + + if not warmup_model: + reconst_loss += ( + ( + loss_z1_weight + + ((loss_z1_unweight).view(self.n_labels_conditioning, -1).t() + * probs_conditioning).sum(dim=1) + ) + * kl_weight + * loss_z1_factor + ) + + if self.prior_z1 == "mog": + cats = Categorical(logits=self.prior_logits) + normal_dists = Independent( + Normal(self.prior_means, torch.exp(self.prior_log_scales) + 1e-4), + 1, + ) + prior = MixtureSameFamily(cats, normal_dists) + u = qz1.rsample(sample_shape=(30,)) + # (sample, n_obs, n_latent) -> (sample, n_obs,) + kl_divergence = -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) + elif self.prior_z1 == "mog_celltype": + if warmup_model: + # Assigns zero meaning equal weight to all unlabeled cells. Otherwise biases to sample from respective MoG. + logits_input = torch.nn.functional.one_hot( + y[:, self.mog_class].ravel().long(), self.num_classes[self.mog_class] + 1 + ).float()[:, :-1] + cats = Categorical(logits=10 * logits_input) + else: + cats = Categorical(logits=logits_conditioning) + normal_dists = torch.distributions.Independent( + Normal( + self.prior_z1_means.expand(x.shape[0], -1, -1), + torch.exp(self.prior_z1_scales).expand(x.shape[0], -1, -1) + 1e-2, + ), + reinterpreted_batch_ndims=1, + ) + + prior = MixtureSameFamily(cats, normal_dists) + u = qz1.rsample(sample_shape=(30,)) + # (sample, n_obs, n_latent) -> (sample, n_obs,) + kl_z = -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) + kl_divergence = weighting_mog * kl_z + else: + prior = Normal(torch.zeros_like(qz1.loc), torch.ones_like(qz1.loc)) + kl_z = 0 + kl_divergence = 0 + + probs_prior, _ = self.classification(tensors, precomputed_z=prior.sample()) + kl_divergence_cat = 0 + + if not warmup_model: + for i in range(0, self.total_level): + y_prior_ = ( + torch.stack( + [self.y_prior_fine[idx] for idx in site_index.ravel().long()], dim=0 + ) + if i == self.total_level - 1 + else getattr(self, f"y_prior_{i}") + ) + if self.update_yprior and i > 0: + if i == self.total_level - 1: + # Shape batch, coarse; batch, coarse, finer -> batch, finer + y_prior_ = torch.einsum("bc,bcf->bf", probs[i - 1], y_prior_) + else: + y_prior_ = torch.einsum("bc,cf->bf", probs[i - 1], y_prior_) + + kl_divergence_cat += kl( + Categorical(probs=probs[i]), + Categorical(probs=y_prior_), + ) + for i in range(0, self.total_level): + if i == self.total_level - 1: + y_prior_ = torch.stack( + [self.y_prior_fine[i] for i in site_index.ravel().long()], dim=0 + ) + else: + y_prior_ = self.__getattr__("y_prior_" + str(i)) + if self.update_yprior and i > 0: + if i == self.total_level - 1: + # Shape batch, coarse; batch, coarse, finer -> batch, finer + y_prior_ = torch.einsum("bc,bcf->bf", probs_prior[i - 1], y_prior_) + else: + y_prior_ = torch.einsum("bc,cf->bf", probs_prior[i - 1], y_prior_) + + kl_divergence_cat += bg_classifier_ratio * kl( + Categorical(probs=probs_prior[i]), + Categorical(probs=y_prior_), + ) + + kl_divergence += kl_divergence_cat + + loss = torch.mean(reconst_loss + kl_divergence * kl_weight) + + if labelled_tensors is not None: + # We filter cells with unlabeled_category in loss. + ce_loss, fine_true_labels, logits_fine = self.classification( + tensors, return_classifier_loss=True + ) + if not warmup_model: + loss += ce_loss * classification_ratio + return LossOutput( + loss=loss, + reconstruction_loss=reconst_loss, + kl_local=kl_divergence, + classification_loss=ce_loss, + true_labels=fine_true_labels, + logits=logits_fine, + ) + return LossOutput(loss=loss, reconstruction_loss=reconst_loss, kl_local=kl_divergence) diff --git a/src/scvi/external/muanvi/_utils.py b/src/scvi/external/muanvi/_utils.py new file mode 100644 index 0000000000..b3aecff88e --- /dev/null +++ b/src/scvi/external/muanvi/_utils.py @@ -0,0 +1,179 @@ +import warnings +from collections.abc import Iterable as IterableClass +from collections.abc import Sequence + +import numpy as np +import pandas as pd +from anndata import AnnData +from pandas.api.types import CategoricalDtype + +from scvi import REGISTRY_KEYS, settings +from scvi.data import AnnDataManager +from scvi.data._utils import _make_column_categorical, get_anndata_attribute +from scvi.data.fields import CategoricalJointObsField +from scvi.dataloaders import ConcatDataLoader +from scvi.dataloaders._ann_dataloader import AnnDataLoader + + +# Class creating a new Obsm field for partially annotated layers of labels +class LabelsWithUnlabeledJointObsField(CategoricalJointObsField): + """ + An AnnDataField for a collection of partially observed layers of labels .obs fields in the AnnData data structure. + + Creates an .obsm field compiling the given .obs fields. The model will reference the compiled + data as a whole. + + Parameters + ---------- + registry_key + Key to register field under in data registry. + attr_keys + Sequence of keys to combine to form the obsm or varm field. + unlabeled_category + A single category to represent unlabeled cells in the data. + """ + + MAPPINGS_KEY = "mappings" + FIELD_KEYS_KEY = "field_keys" + N_CATS_PER_KEY = "n_cats_per_key" + UNLABELED_CATEGORY = "unlabeled_category" + + def __init__( + self, + registry_key: str, + attr_keys: list[str] | None, + unlabeled_category: str | None, + ) -> None: + super().__init__(registry_key, attr_keys) + self.count_stat_key = f"n_{self.registry_key}" + self.unlabeled_category = unlabeled_category + + def _default_mappings_dict(self) -> dict: + return { + self.MAPPINGS_KEY: dict(), + self.FIELD_KEYS_KEY: [], + self.N_CATS_PER_KEY: [], + self.UNLABELED_CATEGORY: [], + } + + def _make_obsm_categorical( + self, adata: AnnData, category_dict: dict[str, list[str]] | None = None + ) -> dict: + if self.attr_keys != getattr(adata, self.attr_name)[self.attr_key].columns.tolist(): + raise ValueError( + f"Original .{self.source_attr_name} keys do not match the columns in the ", + f"generated .{self.attr_name} field.", + ) + + categories = {} + df = getattr(adata, self.attr_name)[self.attr_key] + for level, key in enumerate(self.attr_keys): + categorical_dtype = ( + CategoricalDtype(categories=category_dict[key]) + if category_dict is not None + else None + ) + if categorical_dtype is None: + categorical_obs = df[key].astype("category") + else: + categorical_obs = df[key].astype(categorical_dtype) + + mapping = categorical_obs.cat.categories.to_numpy(copy=True) + mapping = self._remap_unlabeled_to_final_category(mapping, level) + cat_dtype = CategoricalDtype(categories=mapping, ordered=True) + mapping = _make_column_categorical(df, key, key, categorical_dtype=cat_dtype) + categories[key] = mapping + + store_cats = categories if category_dict is None else category_dict + + mappings_dict = self._default_mappings_dict() + mappings_dict[self.MAPPINGS_KEY] = store_cats + mappings_dict[self.FIELD_KEYS_KEY] = self.attr_keys + mappings_dict[self.UNLABELED_CATEGORY] = self.unlabeled_category + for k in self.attr_keys: + mappings_dict[self.N_CATS_PER_KEY].append(len(store_cats[k]) - 1) + return mappings_dict + + def _remap_unlabeled_to_final_category(self, mapping: np.ndarray, level: int) -> np.ndarray: + # Make unlabeled category the last element + unlabeled_category = self.unlabeled_category + + # Check if the unlabeled category is in the mapping + if unlabeled_category in mapping: + # Find the index of the unlabeled category + unlabeled_idx = np.where(mapping == unlabeled_category)[0][0] + # Swap the unlabeled category with the last element + mapping[unlabeled_idx], mapping[-1] = mapping[-1], mapping[unlabeled_idx] + else: + # Append the unlabeled category if it's not in the mapping + mapping = np.append(mapping, unlabeled_category) + + return mapping + + def register_field(self, adata: AnnData) -> dict: + super().register_field(adata) + self._combine_fields(adata) + state_registry = self._make_obsm_categorical(adata) + return state_registry + + def transfer_field( + self, + state_registry: dict, + adata_target: AnnData, + extend_categories: bool = False, + allow_missing_labels: bool = False, + **kwargs, + ) -> dict: + """Transfer the field.""" + for level, key in enumerate(self.attr_keys): + if ( + allow_missing_labels + and key is not None + and key not in list(adata_target.obs.columns) + ): + # Fill in original .obs attribute with unlabeled_category values. + warnings.warn( + f"Missing labels key {key}. Filling in with " + f"unlabeled category {self.unlabeled_category}.", + UserWarning, + stacklevel=settings.warnings_stacklevel, + ) + adata_target.obs[key] = self.unlabeled_category[level] + + kwargs.pop("extend_categories", None) + transfer_state_registry = super().transfer_field( + state_registry, adata_target, extend_categories=extend_categories, **kwargs + ) + categories = {} + mapping = transfer_state_registry[self.MAPPINGS_KEY] + for level, key in enumerate(self.attr_keys): + mapping_ = self._remap_unlabeled_to_final_category(mapping[key], level) + categories[key] = mapping_ + store_cats = categories + mappings_dict = self._default_mappings_dict() + mappings_dict[self.MAPPINGS_KEY] = store_cats + mappings_dict[self.FIELD_KEYS_KEY] = self.attr_keys + mappings_dict[self.UNLABELED_CATEGORY] = self.unlabeled_category + for k in self.attr_keys: + mappings_dict[self.N_CATS_PER_KEY].append(len(store_cats[k]) - 1) + return mappings_dict + + +def _get_site_code_from_category(adata_manager: AnnDataManager, category: Sequence[int | str]): + if not isinstance(category, IterableClass) or isinstance(category, str): + category = [category] + + site_mappings = adata_manager.get_state_registry(REGISTRY_KEYS.SITE_KEY).categorical_mapping + site_code = [] + for cat in category: + if cat is None: + site_code.append(None) + continue + elif isinstance(cat, int) and cat < len(site_mappings): + site_code.append(site_mappings[cat]) + elif cat not in site_mappings: + raise ValueError(f'"{cat}" not a valid site category.') + else: + site_loc = np.where(site_mappings == cat)[0][0] + site_code.append(site_loc) + return site_code, site_mappings diff --git a/tests/external/muanvi/test_muanvi.py b/tests/external/muanvi/test_muanvi.py new file mode 100644 index 0000000000..5d64563866 --- /dev/null +++ b/tests/external/muanvi/test_muanvi.py @@ -0,0 +1,145 @@ +import pytest +from mudata import MuData +import numpy as np +from anndata import AnnData +import itertools + +from scipy import sparse as sp_sparse +import pandas as pd + +from scvi.data import synthetic_iid +from scvi.external import MUANVI + + +# helper function for testing purposes ; could be moved in the same file as generate_synthetic() +def _generate_synthetic_hierarchy( + batch_size: int = 128, + n_genes: int = 100, + n_proteins: int = 100, + n_batches: int = 2, + n_labels_1: int = 2, + n_labels_2: int = 11, + n_sites: int = 2, + sparse: bool = False, +) -> AnnData: + """ + New method to generate test data with two-layer labels. + """ + n_total = batch_size * n_batches * n_sites + data = np.random.negative_binomial(5, 0.3, size=(n_total)) + mask = np.random.binomial(n=1, p=0.7, size=(n_total, n_genes)) + data = data * mask # We put the batch index first + labels_1 = np.random.randint(0, n_labels_1, size=(n_total,)) + labels_1 = np.array([f"label_{i}" for i in labels_1]) + labels_2 = np.random.randint(0, n_labels_2 // n_sites, size=(n_total,)) + + batch = [] + site = [] + labels_2 = [] + for site, batch in itertools.product(range(n_sites), range(n_batches)): + batch += [f"batch_{batch}_site_{site}"] * batch_size + site += [f"site_{site}"] * batch_size + labels_1 = np.array([f"label_{d}_{s}" for d, s in zip(labels_2, site, strict=True)]) + + if sparse: + data = sp_sparse.csr_matrix(data) + adata = AnnData(data) + adata.obs["batch"] = pd.Categorical(batch) + adata.obs["site"] = pd.Categorical(site) + adata.obs["labels_1"] = pd.Categorical(labels_1) + adata.obs["labels_2"] = pd.Categorical(labels_2) + + # Protein measurements + p_data = np.random.negative_binomial(5, 0.3, size=(adata.shape[0], n_proteins)) + adata.obsm["protein_expression"] = p_data + adata.uns["protein_names"] = np.arange(n_proteins).astype(str) + + return adata + + +# helper function for testing purposes ; could be moved int he same file as synthetic_iid() +def synthetic_iid_hierarchy( + batch_size: int | None = 200, + n_genes: int | None = 100, + n_proteins: int | None = 100, + n_batches: int | None = 2, + n_sites: int | None = 2, + n_labels_1: int | None = 2, + n_labels_2: int | None = 10, + sparse: bool = False, +) -> AnnData: + """Synthetic dataset with ZINB distributed RNA and NB distributed protein, with three-layer annotation. + This dataset is just for testing purposed and not meant for modeling or research. + Each value is independently and identically distributed. + Parameters + ---------- + batch_size + Number of cells per batch + n_genes + Number of genes + n_proteins + Number of proteins + n_batches + Number of batches + n_sites + Number of sites + n_labels_1 + Number of cell types of layer 1 + n_labels_2 + Number of cell types of layer 2 + sparse + Whether to use a sparse matrix + Returns + ------- + AnnData with batch info (``.obs['batch']``), label info (``.obs['labels']``), + site info (``.obs['site']``), + protein expression (``.obsm["protein_expression"]``) and + protein names (``.obs['protein_names']``) + Examples + -------- + >>> import scvi + >>> adata = scvi.data.synthetic_iid() + """ + + return _generate_synthetic_hierarchy( + batch_size=batch_size, + n_genes=n_genes, + n_proteins=n_proteins, + n_batches=n_batches, + n_sites=n_sites, + n_labels_1=n_labels_1, + n_labels_2=n_labels_2, + sparse=sparse, + ) + + +def test_methylvi(): + adata1 = synthetic_iid() + adata1.layers["mc"] = adata1.X + adata1.layers["cov"] = adata1.layers["mc"] + 10 + + adata2 = synthetic_iid() + adata2.layers["mc"] = adata2.X + adata2.layers["cov"] = adata2.layers["mc"] + 10 + + mdata = MuData({"mod1": adata1, "mod2": adata2}) + + METHYLVI.setup_mudata( + mdata, + mc_layer="mc", + cov_layer="cov", + methylation_contexts=["mod1", "mod2"], + batch_key="batch", + modalities={"batch_key": "mod1"}, + ) + vae = METHYLVI( + mdata, + ) + vae.train(3) + vae.get_elbo(indices=vae.validation_indices) + vae.get_normalized_methylation() # Retrieve methylation for all contexts + vae.get_normalized_methylation(context="mod1") # Retrieve for specific context + with pytest.raises(ValueError): # Should fail when invalid context selected + vae.get_normalized_methylation(context="mod3") + vae.get_latent_representation() + vae.differential_methylation(groupby="mod1:labels", group1="label_1") From 7e7380e9fef8402e62bf30525e88fa22cce4a40e Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 22 Jan 2025 23:20:55 -0800 Subject: [PATCH 02/24] muANVI --- src/scvi/_constants.py | 1 + src/scvi/data/_utils.py | 3 +- src/scvi/data/fields/_scanvi.py | 1 + src/scvi/external/muanvi/_model.py | 67 ++++----- src/scvi/external/muanvi/_module.py | 33 ++++- src/scvi/external/muanvi/_utils.py | 13 +- src/scvi/model/_scvi.py | 4 + src/scvi/model/base/_archesmixin.py | 2 + src/scvi/model/base/_training_mixin.py | 3 +- src/scvi/module/_constants.py | 2 + src/scvi/module/_vae.py | 191 ++++++++++++++++++++++++- src/scvi/nn/_base_components.py | 83 +++++++++-- src/scvi/train/_trainingplans.py | 11 +- src/scvi/utils/_docstrings.py | 6 + 14 files changed, 347 insertions(+), 73 deletions(-) diff --git a/src/scvi/_constants.py b/src/scvi/_constants.py index 3dc4ea833b..d06c8aabcb 100644 --- a/src/scvi/_constants.py +++ b/src/scvi/_constants.py @@ -6,6 +6,7 @@ class _REGISTRY_KEYS_NT(NamedTuple): ATAC_X_KEY: str = "atac" BATCH_KEY: str = "batch" SITE_KEY: str = "site" + ASSAY_KEY: str = "assay" SAMPLE_KEY: str = "sample" LABELS_KEY: str = "labels" PROTEIN_EXP_KEY: str = "proteins" diff --git a/src/scvi/data/_utils.py b/src/scvi/data/_utils.py index 1b5c5a5253..f1af020e91 100644 --- a/src/scvi/data/_utils.py +++ b/src/scvi/data/_utils.py @@ -197,6 +197,7 @@ def _make_column_categorical( column_key: str, alternate_column_key: str, categorical_dtype: str | CategoricalDtype | None = None, + warning: bool = True, ): """Makes the data in column_key in DataFrame all categorical. @@ -221,7 +222,7 @@ def _make_column_categorical( df[alternate_column_key] = codes # make sure each category contains enough cells - if np.min(counts) < 3: + if np.min(counts) < 3 and warning: category = unique[np.argmin(counts)] warnings.warn( f"Category {category} in adata.obs['{alternate_column_key}'] has fewer than 3 cells. " diff --git a/src/scvi/data/fields/_scanvi.py b/src/scvi/data/fields/_scanvi.py index 6e1d443b5a..db3dcfb85b 100644 --- a/src/scvi/data/fields/_scanvi.py +++ b/src/scvi/data/fields/_scanvi.py @@ -58,6 +58,7 @@ def _remap_unlabeled_to_final_category(self, adata: AnnData, mapping: np.ndarray self._original_attr_key, self.attr_key, categorical_dtype=cat_dtype, + warning=False, ) return { diff --git a/src/scvi/external/muanvi/_model.py b/src/scvi/external/muanvi/_model.py index e2b6e87531..8daa3e0c21 100644 --- a/src/scvi/external/muanvi/_model.py +++ b/src/scvi/external/muanvi/_model.py @@ -110,10 +110,10 @@ def __init__( n_batch = self.summary_stats.n_batch n_fine_labels = self.summary_stats.n_labels - 1 - print(self.summary_stats) n_site = self.summary_stats.n_site + n_assay = self.summary_stats.n_assay - hierarchy_dict, num_classes, hierarchy_matrix = self.extract_hierarchy( + self.hierarchy_dict, self.num_classes, self.hierarchy_matrix = self.extract_hierarchy( n_site=n_site, eps_yprior=eps_yprior ) n_cats_per_cov = ( @@ -128,8 +128,9 @@ def __init__( n_input=self.summary_stats.n_vars, n_batch=n_batch, n_site=n_site, + n_assay=n_assay, n_fine_labels = n_fine_labels, - num_classes=num_classes, + num_classes=self.num_classes, n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), n_cats_per_cov=n_cats_per_cov, n_hidden=n_hidden, @@ -139,8 +140,8 @@ def __init__( dispersion=dispersion, gene_likelihood=gene_likelihood, use_size_factor_key=use_size_factor_key, - hierarchy_dict=hierarchy_dict, - hierarchy_matrix=hierarchy_matrix, + hierarchy_dict=self.hierarchy_dict, + hierarchy_matrix=self.hierarchy_matrix, update_yprior=update_yprior, **muanvae_model_kwargs, ) @@ -252,9 +253,9 @@ def _set_indices_and_labels(self): self._unlabeled_indices = list( set(np.arange(self.adata.n_obs)) - set(self._labeled_indices) ) - self._code_to_label = [ - dict(enumerate(self._label_mapping[layer])) for layer in self._label_mapping - ] + self._code_to_label = { + layer: dict(enumerate(self._label_mapping[layer])) for layer in self._label_mapping + } def extract_hierarchy(self, n_site, eps_yprior): """ @@ -270,31 +271,31 @@ def extract_hierarchy(self, n_site, eps_yprior): labels_state_registry = self.adata_manager.get_state_registry("label_hierarchy") label_keys = labels_state_registry.field_keys - def nested_groupby(df, label_keys): - # Recursively adds dicts for each level of hierarchy. - if len(label_keys) == 1: - return df.groupby(label_keys[0]).apply(list).to_dict() + def fixed_depth_groupby(df, label_keys): + if len(label_keys)==1: + return list(df[label_keys[-1]].unique()) if not df.empty else [] else: + # Recursive case: build dictionaries up to the fixed depth return { - key: nested_groupby(sub_df, label_keys[1:]) - for key, sub_df in df.groupby(label_keys[0]) + key: fixed_depth_groupby(sub_df, label_keys[1:]) + for key, sub_df in df.groupby(label_keys[0], observed=True) } - hierarchy_dict = nested_groupby(self.labels, label_keys) num_classes = labels_state_registry.n_cats_per_key + hierarchy_dict = fixed_depth_groupby(self.labels, label_keys) hierarchy_matrix = [] for n_label in range(1, len(num_classes)): + curr = pd.DataFrame( + 0, + index=np.arange(num_classes[n_label - 1]), + columns=np.arange(num_classes[n_label]), + ) # Site specific last layer. if n_label == len(num_classes) - 1: hierarchy_matrix_ = torch.zeros( num_classes[n_label - 1], num_classes[n_label], n_site ) - curr = pd.DataFrame( - 0, - index=np.arange(hierarchy_matrix_.shape[0]), - columns=np.arange(hierarchy_matrix_.shape[1]), - ) for site in range(n_site): adata = self.adata[self.adata.obs["_scvi_site"] == site] @@ -307,8 +308,8 @@ def nested_groupby(df, label_keys): ) curr_[curr_ > 0] = 1 curr_ = curr_.loc[curr.index, curr.columns] - curr = curr_.div(curr_.sum(axis=1), axis=0) - hierarchy_matrix_[:, :, site] = torch.tensor(curr.fillna(0).values) + curr = curr_.div(curr_.sum(axis=1), axis=0).fillna(0) + hierarchy_matrix_[:, :, site] = torch.tensor(curr.values) else: curr_ = pd.crosstab( self.adata.obsm["_scvi_label_hierarchy"][label_keys[n_label - 1]], @@ -316,8 +317,8 @@ def nested_groupby(df, label_keys): ) curr_[curr_ > 0] = 1 curr_ = curr_.loc[curr.index, curr.columns] - curr = curr_[curr_.index, curr_.columns].div(curr_.sum(axis=1), axis=0).values - hierarchy_matrix_ = torch.tensor(curr.fillna(0).values) + curr = curr_.div(curr_.sum(axis=1), axis=0).fillna(0) + hierarchy_matrix_ = torch.tensor(curr.values) hierarchy_matrix.append(hierarchy_matrix_) @@ -378,7 +379,6 @@ def predict( for _, tensors in enumerate(scdl): for site_to_predict_ in sites_to_predict_: probs_, _ = self.module.classification(tensors, site_to_predict=site_to_predict_) - if site_to_predict_ is not None: if not soft: pred_ = probs_.argmax(dim=1) @@ -392,19 +392,14 @@ def predict( for i in range(total_level): if not soft: - probs[i] = probs[i].argmax(dim=1) # select only the label + pred[class_labels[i]].append(probs[i].argmax(dim=1).cpu()) else: pred[class_labels[i]].append(probs[i].detach().cpu()) for key in pred.keys(): pred[key] = torch.cat(pred[key]).numpy() if not soft: - if key not in class_labels: - pred[key] = [self._code_to_label[-1][ct] for ct in pred[key]] - else: - pred[key] = [ - self._code_to_label[class_labels.index(key)][ct] for ct in pred[key] - ] + pred[key] = [self._code_to_label[key][ct] for ct in pred[key]] if not soft: pred = pd.DataFrame.from_dict(pred) @@ -412,10 +407,7 @@ def predict( return pred else: for key in pred.keys(): - if key not in class_labels: - columns = list(self._code_to_label[-1].values())[:-1] - else: - columns = list(self._code_to_label[class_labels.index(key)].values())[:-1] + columns = list(self._code_to_label[key].values())[:-1] pred[key] = pd.DataFrame( pred[key], @@ -684,6 +676,7 @@ def setup_anndata( unlabeled_category: list[str | int | float], layer: str | None = None, site_key: str | None = None, + assay_key: str | None = None, batch_key: str | None = None, size_factor_key: str | None = None, categorical_covariate_keys: list[str] | None = None, @@ -699,6 +692,7 @@ def setup_anndata( %(param_layer)s %(param_batch_key)s %(param_site_key)s + %(param_assay_key)s fine_labels_key key in `adata.obs` for fine label information. Categories will automatically be converted into integer categories and saved to `adata.obs['_scvi_labels']`. @@ -717,6 +711,7 @@ def setup_anndata( anndata_fields = [ LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), + CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), CategoricalObsField(REGISTRY_KEYS.SITE_KEY, site_key), LabelsWithUnlabeledObsField(REGISTRY_KEYS.LABELS_KEY, fine_labels_key, unlabeled_category), NumericalObsField(REGISTRY_KEYS.SIZE_FACTOR_KEY, size_factor_key, required=False), diff --git a/src/scvi/external/muanvi/_module.py b/src/scvi/external/muanvi/_module.py index 3fe1b5a30e..688cd6ed0d 100644 --- a/src/scvi/external/muanvi/_module.py +++ b/src/scvi/external/muanvi/_module.py @@ -26,6 +26,10 @@ class MUANVAE(SCANVAE): Number of input genes n_batch Number of batches + n_site + Number of annotation sites + n_assay + Number of assays n_fine_labels Number of fine labels num_classes @@ -75,6 +79,7 @@ def __init__( hierarchy_dict: dict, n_batch: int = 0, n_site: int = 0, + n_assay: int = 0, n_fine_labels: int = 0, n_hidden: int = 128, n_latent: int = 10, @@ -85,8 +90,8 @@ def __init__( dispersion: str = "gene", log_variational: bool = True, gene_likelihood: str = "nb", - classifier_parameters: dict = dict(), - classifier_parameters_muanvae: dict = dict(), + classifier_parameters: dict | None= None, + classifier_parameters_muanvae: dict | None= None, use_batch_norm: Literal["encoder", "decoder", "none", "both"] = "none", use_layer_norm: Literal["encoder", "decoder", "none", "both"] = "both", conditioning_class: int = -1, @@ -101,13 +106,17 @@ def __init__( self.mog_class = mog_class self.site_specific_classifier = n_site > 1 self.n_site = n_site + self.n_assay = n_assay self.num_classes = num_classes self.update_yprior = update_yprior self.n_labels_conditioning = num_classes[self.conditioning_class] self.n_fine_labels = n_fine_labels self.hiearchy_dict = hierarchy_dict - classifier_parameters = classifier_parameters or {} + if classifier_parameters is None: + classifier_parameters = {} + if classifier_parameters_muanvae is None: + classifier_parameters_muanvae = {} cls_parameters = { "n_layers": n_layers, @@ -344,6 +353,19 @@ def loss( logits[self.conditioning_class], ) + if z1.ndim == 2: + loss_z1_unweight_ = loss_z1_unweight.view(self.n_labels, -1).t() + kl_divergence_z2_ = kl_divergence_z2.view(self.n_labels, -1).t() + else: + loss_z1_unweight_ = torch.transpose( + loss_z1_unweight.view(z1.shape[0], self.n_labels, -1), -1, -2 + ) + kl_divergence_z2_ = torch.transpose( + kl_divergence_z2.view(z1.shape[0], self.n_labels, -1), -1, -2 + ) + reconst_loss += loss_z1_weight + (loss_z1_unweight_ * probs[-1]).sum(dim=-1) + kl_divergence = (kl_divergence_z2_ * probs[-1]).sum(dim=-1) + if not warmup_model: reconst_loss += ( ( @@ -364,7 +386,7 @@ def loss( prior = MixtureSameFamily(cats, normal_dists) u = qz1.rsample(sample_shape=(30,)) # (sample, n_obs, n_latent) -> (sample, n_obs,) - kl_divergence = -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) + kl_divergence += -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) elif self.prior_z1 == "mog_celltype": if warmup_model: # Assigns zero meaning equal weight to all unlabeled cells. Otherwise biases to sample from respective MoG. @@ -386,11 +408,10 @@ def loss( u = qz1.rsample(sample_shape=(30,)) # (sample, n_obs, n_latent) -> (sample, n_obs,) kl_z = -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) - kl_divergence = weighting_mog * kl_z + kl_divergence += weighting_mog * kl_z else: prior = Normal(torch.zeros_like(qz1.loc), torch.ones_like(qz1.loc)) kl_z = 0 - kl_divergence = 0 probs_prior, _ = self.classification(tensors, precomputed_z=prior.sample()) kl_divergence_cat = 0 diff --git a/src/scvi/external/muanvi/_utils.py b/src/scvi/external/muanvi/_utils.py index b3aecff88e..a93b9f12a3 100644 --- a/src/scvi/external/muanvi/_utils.py +++ b/src/scvi/external/muanvi/_utils.py @@ -3,16 +3,13 @@ from collections.abc import Sequence import numpy as np -import pandas as pd from anndata import AnnData from pandas.api.types import CategoricalDtype from scvi import REGISTRY_KEYS, settings from scvi.data import AnnDataManager -from scvi.data._utils import _make_column_categorical, get_anndata_attribute +from scvi.data._utils import _make_column_categorical from scvi.data.fields import CategoricalJointObsField -from scvi.dataloaders import ConcatDataLoader -from scvi.dataloaders._ann_dataloader import AnnDataLoader # Class creating a new Obsm field for partially annotated layers of labels @@ -81,7 +78,13 @@ def _make_obsm_categorical( mapping = categorical_obs.cat.categories.to_numpy(copy=True) mapping = self._remap_unlabeled_to_final_category(mapping, level) cat_dtype = CategoricalDtype(categories=mapping, ordered=True) - mapping = _make_column_categorical(df, key, key, categorical_dtype=cat_dtype) + mapping = _make_column_categorical( + df, + key, + key, + categorical_dtype=cat_dtype, + warning=False + ) categories[key] = mapping store_cats = categories if category_dict is None else category_dict diff --git a/src/scvi/model/_scvi.py b/src/scvi/model/_scvi.py index e596883e07..fb9ee4e59b 100644 --- a/src/scvi/model/_scvi.py +++ b/src/scvi/model/_scvi.py @@ -180,6 +180,7 @@ def __init__( n_cats_per_cov = None n_batch = self.summary_stats.n_batch + n_assay = self.summary_stats.n_assay use_size_factor_key = self.registry_["setup_args"][ f"{REGISTRY_KEYS.SIZE_FACTOR_KEY}_key" ] @@ -195,6 +196,7 @@ def __init__( self.module = self._module_cls( n_input=self.summary_stats.n_vars, n_batch=n_batch, + n_assay=n_assay, n_labels=self.summary_stats.n_labels, n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), n_cats_per_cov=n_cats_per_cov, @@ -222,6 +224,7 @@ def setup_anndata( adata: AnnData, layer: str | None = None, batch_key: str | None = None, + assay_key: str | None = None, labels_key: str | None = None, size_factor_key: str | None = None, categorical_covariate_keys: list[str] | None = None, @@ -244,6 +247,7 @@ def setup_anndata( anndata_fields = [ LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), + CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), CategoricalObsField(REGISTRY_KEYS.LABELS_KEY, labels_key), NumericalObsField(REGISTRY_KEYS.SIZE_FACTOR_KEY, size_factor_key, required=False), CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), diff --git a/src/scvi/model/base/_archesmixin.py b/src/scvi/model/base/_archesmixin.py index 25e1b6e1fb..64222ea206 100644 --- a/src/scvi/model/base/_archesmixin.py +++ b/src/scvi/model/base/_archesmixin.py @@ -231,6 +231,7 @@ def load_query_data( freeze_dropout=freeze_dropout, freeze_expression=freeze_expression, freeze_classifier=freeze_classifier, + parameters_yes_grad=additional_parameters, ) model.is_trained_ = False @@ -413,6 +414,7 @@ def _set_params_online_update( if not freeze_classifier: mod_no_hooks_yes_grad.add("classifier") parameters_yes_grad = {"background_pro_alpha", "background_pro_log_beta"} + parameters_yes_grad = {"background_pro_alpha", "background_pro_log_beta"} def no_hook_cond(key): one = (not freeze_expression) and "encoder" in key diff --git a/src/scvi/model/base/_training_mixin.py b/src/scvi/model/base/_training_mixin.py index cc7d935891..85b2671669 100644 --- a/src/scvi/model/base/_training_mixin.py +++ b/src/scvi/model/base/_training_mixin.py @@ -14,6 +14,7 @@ from scvi.dataloaders import DataSplitter, SemiSupervisedDataSplitter from scvi.model._utils import get_max_epochs_heuristic, use_distributed_sampler from scvi.train import ( + AdversarialTrainingPlan, SemiSupervisedAdversarialTrainingPlan, SemiSupervisedTrainingPlan, TrainingPlan, @@ -40,7 +41,7 @@ class UnsupervisedTrainingMixin: """General purpose unsupervised train method.""" _data_splitter_cls = DataSplitter - _training_plan_cls = TrainingPlan + _training_plan_cls = AdversarialTrainingPlan _train_runner_cls = TrainRunner @devices_dsp.dedent diff --git a/src/scvi/module/_constants.py b/src/scvi/module/_constants.py index 2b6e232429..f4fabc4ac1 100644 --- a/src/scvi/module/_constants.py +++ b/src/scvi/module/_constants.py @@ -11,6 +11,8 @@ class _MODULE_KEYS(NamedTuple): LIBRARY_KEY: str = "library" QL_KEY: str = "ql" BATCH_INDEX_KEY: str = "batch_index" + ASSAY_INDEX_KEY: str = "assay_index" + SITE_INDEX_KEY: str = "site_index" Y_KEY: str = "y" CONT_COVS_KEY: str = "cont_covs" CAT_COVS_KEY: str = "cat_covs" diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index 183ccb5b5a..e9277d0986 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -7,6 +7,7 @@ import numpy as np import torch from torch.nn.functional import one_hot +from torch import distributions from scvi import REGISTRY_KEYS, settings from scvi.data._constants import ADATA_MINIFY_TYPE @@ -147,6 +148,7 @@ def __init__( self, n_input: int, n_batch: int = 0, + n_assay: int = 0, n_labels: int = 0, n_hidden: int = 128, n_latent: int = 10, @@ -172,6 +174,10 @@ def __init__( extra_encoder_kwargs: dict | None = None, extra_decoder_kwargs: dict | None = None, batch_embedding_kwargs: dict | None = None, + conditional_norm: dict | None = None, + mmd_kernel: str | None = "rbf", + prior: str | None = None, + num_classes: int | None = 30, ): from scvi.nn import DecoderSCVI, Encoder @@ -183,6 +189,7 @@ def __init__( self.gene_likelihood = gene_likelihood self.n_batch = n_batch self.n_input = n_input + self.n_assay = n_assay self.n_labels = n_labels self.n_hidden = n_hidden self.n_layers = n_layers @@ -191,6 +198,7 @@ def __init__( 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 + self.mmd_kernel = mmd_kernel if not self.use_observed_lib_size: if library_log_means is None or library_log_vars is None: @@ -233,8 +241,13 @@ def __init__( 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) + if n_assay > 1: + encoder_cat_list = [n_assay] + list([] if n_cats_per_cov is None else n_cats_per_cov) + n_input_encoder = n_input + n_continuous_cov * encode_covariates + else: + encoder_cat_list = cat_list - encoder_cat_list = cat_list if encode_covariates else None + encoder_cat_list = encoder_cat_list if encode_covariates else None _extra_encoder_kwargs = extra_encoder_kwargs or {} self.z_encoder = Encoder( n_input_encoder, @@ -249,6 +262,7 @@ def __init__( use_layer_norm=use_layer_norm_encoder, var_activation=var_activation, return_dist=True, + conditional_norm=conditional_norm, **_extra_encoder_kwargs, ) # l encoder goes from n_input-dimensional data to 1-d library size @@ -277,6 +291,7 @@ def __init__( n_cat_list=cat_list, n_layers=n_layers, n_hidden=n_hidden, + n_assay=self.n_assay, inject_covariates=deeply_inject_covariates, use_batch_norm=use_batch_norm_decoder, use_layer_norm=use_layer_norm_decoder, @@ -284,6 +299,20 @@ def __init__( **_extra_decoder_kwargs, ) + self.prior = prior + if prior == "mog": + self.register_parameter( + "prior_means", + torch.nn.Parameter(torch.randn([num_classes, n_latent])), + ) + self.register_parameter( + "prior_log_scales", + torch.nn.Parameter(torch.zeros([num_classes, n_latent])), + ) + self.register_parameter( + "prior_logits", torch.nn.Parameter(torch.ones([num_classes])) + ) + def _get_inference_input( self, tensors: dict[str, torch.Tensor | None], @@ -304,6 +333,8 @@ def _get_inference_input( return { MODULE_KEYS.X_KEY: tensors[REGISTRY_KEYS.X_KEY], MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY], + MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY], + MODULE_KEYS.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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), } @@ -328,6 +359,7 @@ def _get_generative_input( 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.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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), @@ -362,6 +394,7 @@ def _regular_inference( self, x: torch.Tensor, batch_index: torch.Tensor, + assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, n_samples: int = 1, @@ -371,6 +404,7 @@ def _regular_inference( if self.use_observed_lib_size: library = torch.log(x.sum(1)).unsqueeze(1) if self.log_variational: + x_ = x_/x_.mean(1).unsqueeze(1) x_ = torch.log1p(x_) if cont_covs is not None and self.encode_covariates: @@ -382,10 +416,13 @@ def _regular_inference( else: categorical_input = () - if self.batch_representation == "embedding" and self.encode_covariates: + if assay_index is not None: + assay_index = assay_index.long() + qz, z = self.z_encoder(encoder_input, assay_index, *categorical_input) + elif self.encode_covariates and self.batch_representation == "embedding": 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) + qz, z = self.z_encoder(encoder_input, batch_index, *categorical_input) else: qz, z = self.z_encoder(encoder_input, batch_index, *categorical_input) @@ -448,6 +485,7 @@ def generative( z: torch.Tensor, library: torch.Tensor, batch_index: torch.Tensor, + assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, size_factor: torch.Tensor | None = None, @@ -495,6 +533,7 @@ def generative( size_factor, *categorical_input, y, + assay=assay_index.long() if assay_index is not None else None, ) else: px_scale, px_r, px_rate, px_dropout = self.decoder( @@ -504,6 +543,7 @@ def generative( batch_index, *categorical_input, y, + assay=assay_index.long() if assay_index is not None else None, ) if self.dispersion == "gene-label": @@ -555,14 +595,27 @@ def loss( inference_outputs: dict[str, torch.Tensor | Distribution | None], generative_outputs: dict[str, Distribution | None], kl_weight: torch.tensor | float = 1.0, + weight_assay_loss: 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 self.prior == "mog": + qz = inference_outputs[MODULE_KEYS.QZ_KEY] + cats = distributions.Categorical(logits=self.prior_logits) + normal_dists = distributions.Independent( + distributions.Normal(self.prior_means, torch.exp(self.prior_log_scales) + 1e-4), + 1, + ) + prior = distributions.MixtureSameFamily(cats, normal_dists) + u = qz.rsample(sample_shape=(30,)) + # (sample, n_obs, n_latent) -> (sample, n_obs,) + kl_divergence_z = (qz.log_prob(u).sum(-1) - prior.log_prob(u)).mean(0) + else: + 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] @@ -570,6 +623,15 @@ def loss( else: kl_divergence_l = torch.zeros_like(kl_divergence_z) + assay_index = tensors.get(REGISTRY_KEYS.ASSAY_KEY, None) + if weight_assay_loss > 0.0 and assay_index is not None: + assay_loss = self._compute_assay_penalty( + inference_outputs[MODULE_KEYS.QZ_KEY].loc, + assay_index + ) + else: + assay_loss = 0.0 + reconst_loss = -generative_outputs[MODULE_KEYS.PX_KEY].log_prob(x).sum(-1) kl_local_for_warmup = kl_divergence_z @@ -577,7 +639,7 @@ def loss( weighted_kl_local = kl_weight * kl_local_for_warmup + kl_local_no_warmup - loss = torch.mean(reconst_loss + weighted_kl_local) + loss = torch.mean(reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss) # a payload to be used during autotune if self.extra_payload_autotune: @@ -746,6 +808,121 @@ def marginal_ll( batch_log_lkl = batch_log_lkl.cpu() return batch_log_lkl + def _compute_assay_penalty( + self, params, assay): + assay = assay.ravel().long() + unique = torch.unique(assay) + pair_penalty = torch.tensor(0., device=assay.device) + if len(unique) > 1: + for i in unique: + pp = self.mmd(params, mask=(assay == i)) + pair_penalty += pp + + return pair_penalty + + def mmd(self, params, mask=None): + if mask is not None: + mod_1 = params[mask] + mod_2 = params[~mask] + if self.mmd_kernel == 'imq': + penalty = imq_kernel(mod_1, mod_2, beta=0.5) + elif self.mmd_kernel == 'rbf': + penalty = rbf_kernel(mod_1, mod_2) + else: + penalty = torch.linalg.norm(mod_1 - mod_2, dim=1).mean() + return penalty + +def imq_kernel(x, y, gammas=None, beta=0.5): + """ + Compute the IMQ kernel between two tensors. + + Parameters + ---------- + x : torch.Tensor + Input tensor of shape (N, D). + y : torch.Tensor + Input tensor of shape (M, D). + gammas : list of float or None + List of gamma values to compute the kernel with. + beta : float + The beta parameter controlling the sharpness of the kernel. + + Returns + ------- + kernel_sum : torch.Tensor + Tensor of shape (N, M), the sum of IMQ kernels over all gamma values. + """ + if gammas is None: + gammas = [ + 1e-3, + 1e-2, + 1e-1, + 1, + 5, + 10, + ] + kxy = torch.cdist(x, y).pow(2) + kxx = torch.cdist(x, x).pow(2) + kyy = torch.cdist(y, y).pow(2) + + kernel_sum_xy = torch.tensor(0.0, device=x.device) + kernel_sum_xx = torch.tensor(0.0, device=x.device) + kernel_sum_yy = torch.tensor(0.0, device=x.device) + for gamma in gammas: + kernel_sum_xy += (gamma + kxy).pow(-beta).mean() + kernel_sum_xx += (gamma + kxx).pow(-beta).mean() + kernel_sum_yy += (gamma + kyy).pow(-beta).mean() + + return (kernel_sum_xx + kernel_sum_yy - 2 * kernel_sum_xy) / len(gammas) + +@auto_move_data +def rbf_kernel(x, y, gammas=None): + """ + Compute the RBF kernel between two tensors. + + Parameters + ---------- + x : torch.Tensor + Input tensor of shape (N, D). + y : torch.Tensor + Input tensor of shape (M, D). + gammas : list of float or None + List of gamma values to compute the kernel with. + + Returns + ------- + kernel_sum : torch.Tensor + Tensor of shape (N, M), the sum of RBF kernels over all gamma values. + """ + if gammas is None: + gammas = [ + 1e-10, + 1e-8, + 1e-6, + 1e-4, + 1e-3, + 1e-2, + 1e-1, + 1, + 2, + 5, + 10, + ] + kxy = torch.cdist(x, y).pow(2) + kxx = torch.cdist(x, x).pow(2) + kyy = torch.cdist(y, y).pow(2) + + kernel_sum_xy = torch.tensor(0.0, device=x.device) + kernel_sum_xx = torch.tensor(0.0, device=x.device) + kernel_sum_yy = torch.tensor(0.0, device=x.device) + for gamma in gammas: + kernel_sum_xy += torch.exp(-gamma * kxy).mean() + kernel_sum_xx += torch.exp(-gamma * kxx).mean() + kernel_sum_yy += torch.exp(-gamma * kyy).mean() + #print(gamma, (torch.exp(-gamma * kxx).mean() + torch.exp(-gamma * kyy).mean() - 2 * torch.exp(-gamma * kxy).mean())) + + return (kernel_sum_xx + kernel_sum_yy - 2 * kernel_sum_xy) / len(gammas) + class LDVAE(VAE): """Linear-decoded Variational auto-encoder model. diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index fa3e206230..b328c26b6f 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -14,6 +14,38 @@ def _identity(x): return x +class ConditionalBatchNorm2d(nn.Module): + def __init__(self, num_features, num_classes, momentum, eps): + super().__init__() + self.num_features = num_features + self.bn = nn.BatchNorm1d(self.num_features, momentum=momentum, eps=eps, affine=False) + self.embed = nn.Embedding(num_classes, self.num_features * 2) + self.embed.weight.data[:, :self.num_features].normal_(1, 0.02) # Initialise scale at N(1, 0.02) + self.embed.weight.data[:, self.num_features:].zero_() # Initialise bias at 0 + + def forward(self, x, y): + out = self.bn(x) + gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) + out = gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) + + return out + +class ConditionalLayerNorm(nn.Module): + def __init__(self, num_features, num_classes): + super().__init__() + self.num_features = num_features + self.ln = nn.LayerNorm(self.num_features, elementwise_affine=False) + self.embed = nn.Embedding(num_classes, self.num_features * 2) + self.embed.weight.data[:, :self.num_features].normal_(1, 0.02) # Initialise scale at N(1, 0.02) + self.embed.weight.data[:, self.num_features:].zero_() # Initialise bias at 0 + + def forward(self, x, y): + out = self.ln(x) + gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) + out = gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) + + return out + class FCLayers(nn.Module): """A helper class to build fully-connected layers for a neural network. @@ -64,6 +96,7 @@ def __init__( bias: bool = True, inject_covariates: bool = True, activation_fn: nn.Module = nn.ReLU, + conditional_norm: bool = False, ): super().__init__() self.inject_covariates = inject_covariates @@ -90,11 +123,14 @@ def __init__( ), # non-default params come from defaults in the original Tensorflow # implementation - nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) - if use_batch_norm + ConditionalBatchNorm2d(n_out, self.n_cat_list[0], momentum=0.01, eps=0.001) + if conditional_norm and use_batch_norm + else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) if use_batch_norm else None, - nn.LayerNorm(n_out, elementwise_affine=False) - if use_layer_norm + # non-default params come from defaults in original Tensorflow implementation + ConditionalLayerNorm(n_out, self.n_cat_list[0]) + if conditional_norm and use_layer_norm + else nn.LayerNorm(n_out, elementwise_affine=False) if use_layer_norm else None, activation_fn() if use_activation else None, nn.Dropout(p=dropout_rate) if dropout_rate > 0 else None, @@ -175,7 +211,15 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = N for i, layers in enumerate(self.fc_layers): for layer in layers: if layer is not None: - if isinstance(layer, nn.BatchNorm1d): + if isinstance(layer, ConditionalBatchNorm2d) or isinstance( + layer, ConditionalLayerNorm): + if x.dim() == 3: + x = torch.cat( + [(layer(slice_x, cat_list[0])).unsqueeze(0) for slice_x in x], dim=0 + ) + else: + x = layer(x=x, y=cat_list[0]) + elif isinstance(layer, nn.BatchNorm1d): if x.dim() == 3: if ( x.device.type == "mps" @@ -349,10 +393,11 @@ def __init__( n_cat_list: Iterable[int] = None, n_layers: int = 1, n_hidden: int = 128, + n_assay: int = 1, inject_covariates: bool = True, use_batch_norm: bool = False, use_layer_norm: bool = False, - scale_activation: Literal["softmax", "softplus"] = "softmax", + scale_activation: Literal["softmax", "softplus", "exp"] = "softmax", **kwargs, ): super().__init__() @@ -374,16 +419,20 @@ def __init__( px_scale_activation = nn.Softmax(dim=-1) elif scale_activation == "softplus": px_scale_activation = nn.Softplus() + elif scale_activation == "exp": + px_scale_activation = ExpActivation() self.px_scale_decoder = nn.Sequential( - nn.Linear(n_hidden, n_output), + nn.Linear(n_hidden + n_assay, n_output), px_scale_activation, ) - # dispersion: here we only deal with a gene-cell dispersion case - self.px_r_decoder = nn.Linear(n_hidden, n_output) + # dispersion: here we only deal with gene-cell dispersion case + self.px_r_decoder = nn.Linear(n_hidden + n_assay, n_output) # dropout - self.px_dropout_decoder = nn.Linear(n_hidden, n_output) + self.px_dropout_decoder = nn.Linear(n_hidden + n_assay, n_output) + + self.n_assay = n_assay def forward( self, @@ -391,6 +440,7 @@ def forward( z: torch.Tensor, library: torch.Tensor, *cat_list: int, + assay: torch.Tensor | None = None, ): """The forward computation for a single sample. @@ -407,6 +457,8 @@ def forward( * ``'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 + assay + tensor with shape ``(n_input,)`` of assay column z tensor with shape ``(n_input,)`` library @@ -422,11 +474,16 @@ def forward( """ # The decoder returns values for the parameters of the ZINB distribution px = self.px_decoder(z, *cat_list) - px_scale = self.px_scale_decoder(px) - px_dropout = self.px_dropout_decoder(px) + if assay is not None: + one_hot_cat = nn.functional.one_hot(assay.squeeze(-1), self.n_assay) + else: + one_hot_cat = torch.zeros(px.size(0), self.n_assay) + px_cat = torch.cat([px, one_hot_cat], dim=-1) + px_scale = self.px_scale_decoder(px_cat) + px_dropout = self.px_dropout_decoder(px_cat) # Clamp to high value: exp(12) ~ 160000 to avoid nans (computational stability) px_rate = torch.exp(library) * px_scale # torch.clamp( , max=12) - px_r = self.px_r_decoder(px) if dispersion == "gene-cell" else None + px_r = self.px_r_decoder(px_cat) if dispersion == "gene-cell" else None return px_scale, px_r, px_rate, px_dropout diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 8373d9d41f..aef615bd4b 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -635,14 +635,17 @@ def __init__( ) self.adversarial_classifier = False else: - self.n_output_classifier = self.module.n_batch + self.n_output_classifier = self.module.n_assay self.adversarial_classifier = Classifier( n_input=self.module.n_latent, - n_hidden=32, + n_hidden=128, n_labels=self.n_output_classifier, - n_layers=2, + n_layers=1, logits=True, + use_batch_norm=False, + use_layer_norm=True, ) + print('new classifier') else: self.adversarial_classifier = adversarial_classifier self.scale_adversarial_loss = scale_adversarial_loss @@ -676,7 +679,7 @@ def training_step(self, batch, batch_idx): if self.scale_adversarial_loss == "auto" else self.scale_adversarial_loss ) - batch_tensor = batch[REGISTRY_KEYS.BATCH_KEY] + batch_tensor = batch[REGISTRY_KEYS.ASSAY_KEY].long() opts = self.optimizers() if not isinstance(opts, list): diff --git a/src/scvi/utils/_docstrings.py b/src/scvi/utils/_docstrings.py index 1963cab7cb..84d08e6542 100644 --- a/src/scvi/utils/_docstrings.py +++ b/src/scvi/utils/_docstrings.py @@ -124,6 +124,12 @@ integer categories and saved to `adata.obs['_scvi_batch']`. If `None`, assigns the same batch to all the data.""" +param_assay_key = """\ +assay_key + key in `adata.obs` for assay and suspension type information. Categories will automatically be + converted into integer categories and saved to `adata.obs['_scvi_assay']`. If `None`, assigns + the same batch to all the data.""" + param_sample_key = """\ sample_key key in `adata.obs` for sample information. Categories will automatically be converted into From b1b243a185af1e3cd25b64165b01b7a3646fcbd3 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Mon, 24 Feb 2025 00:02:01 -0800 Subject: [PATCH 03/24] Implementation assayvi --- src/scvi/external/__init__.py | 5 +- src/scvi/external/assayvi/__init__.py | 4 + src/scvi/external/assayvi/_model.py | 363 ++++++++++++ src/scvi/external/assayvi/_module.py | 712 +++++++++++++++++++++++ src/scvi/model/base/_embedding_mixin.py | 28 +- src/scvi/module/_constants.py | 1 + src/scvi/module/_vae.py | 13 +- src/scvi/module/base/_embedding_mixin.py | 68 ++- src/scvi/module/base/_priors.py | 2 +- src/scvi/nn/_base_components.py | 100 +++- src/scvi/nn/_embedding.py | 8 +- src/scvi/train/_trainingplans.py | 63 +- 12 files changed, 1288 insertions(+), 79 deletions(-) create mode 100644 src/scvi/external/assayvi/__init__.py create mode 100644 src/scvi/external/assayvi/_model.py create mode 100644 src/scvi/external/assayvi/_module.py diff --git a/src/scvi/external/__init__.py b/src/scvi/external/__init__.py index 4f7d67018c..558af1ccdb 100644 --- a/src/scvi/external/__init__.py +++ b/src/scvi/external/__init__.py @@ -3,6 +3,7 @@ from scvi import settings from scvi.utils import error_on_missing_dependencies +from .assayvi import ASSAYVI from .cellassign import CellAssign from .contrastivevi import ContrastiveVI from .cytovi import CYTOVI @@ -11,11 +12,8 @@ from .gimvi import GIMVI from .methylvi import METHYLANVI, METHYLVI from .mrvi import MRVI -<<<<<<< HEAD from .mrvi_torch import TorchMRVI -======= from .muanvi import MUANVI ->>>>>>> bb720b5f (Add muANVI) from .poissonvi import POISSONVI from .resolvi import RESOLVI from .scar import SCAR @@ -28,6 +26,7 @@ from .velovi import VELOVI __all__ = [ + "ASSAYVI", "SCAR", "SOLO", "GIMVI", diff --git a/src/scvi/external/assayvi/__init__.py b/src/scvi/external/assayvi/__init__.py new file mode 100644 index 0000000000..e7719ae44e --- /dev/null +++ b/src/scvi/external/assayvi/__init__.py @@ -0,0 +1,4 @@ +from ._model import ASSAYVI +from ._module import ASSAYVAE + +__all__ = ["ASSAYVI", "ASSAYVAE"] diff --git a/src/scvi/external/assayvi/_model.py b/src/scvi/external/assayvi/_model.py new file mode 100644 index 0000000000..2e8d0fa261 --- /dev/null +++ b/src/scvi/external/assayvi/_model.py @@ -0,0 +1,363 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +import numpy as np + +from scvi import REGISTRY_KEYS +from scvi.data import AnnDataManager +from scvi.data._utils import _get_adata_minify_type +from scvi.data.fields import ( + CategoricalJointObsField, + CategoricalObsField, + LayerField, + NumericalJointObsField, +) +from scvi.dataloaders import DataSplitter +from scvi.model._utils import ( + get_max_epochs_heuristic, +) +from scvi.model.base import ( + ArchesMixin, + BaseMinifiedModeModelClass, + EmbeddingMixin, + RNASeqMixin, + VAEMixin, +) +from scvi.train import AdversarialTrainingPlan, TrainRunner +from scvi.utils import setup_anndata_dsp +from scvi.utils._docstrings import devices_dsp, setup_anndata_dsp + +from ._module import ASSAYVAE + +if TYPE_CHECKING: + from typing import Literal + + import numpy as np + from anndata import AnnData + +logger = logging.getLogger(__name__) + + +class ASSAYVI( + EmbeddingMixin, + RNASeqMixin, + VAEMixin, + ArchesMixin, + BaseMinifiedModeModelClass, +): + """single-cell Variational Inference :cite:p:`Lopez18`. + + Parameters + ---------- + adata + AnnData object that has been registered via :meth:`~scvi.model.SCVI.setup_anndata`. If + ``None``, then the underlying module will not be initialized until training, and a + :class:`~lightning.pytorch.core.LightningDataModule` must be passed in during training + (``EXPERIMENTAL``). + n_hidden + Number of nodes per hidden layer. + n_latent + Dimensionality of the latent space. + n_layers + Number of hidden layers used for encoder and decoder 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-cell'`` - dispersion can differ for every gene in every cell + gene_likelihood + One of: + + * ``'nb'`` - Negative binomial distribution + * ``'zinb'`` - Zero-inflated negative binomial distribution + * ``'poisson'`` - Poisson distribution + * ``'normal'`` - ``EXPERIMENTAL`` Normal distribution + latent_distribution + One of: + + * ``'normal'`` - Normal distribution + * ``'ln'`` - Logistic normal distribution (Normal(0, I) transformed by softmax) + **kwargs + Additional keyword arguments for :class:`~scvi.module.VAE`. + + Examples + -------- + >>> adata = anndata.read_h5ad(path_to_anndata) + >>> scvi.model.SCVI.setup_anndata(adata, batch_key="batch") + >>> vae = scvi.model.SCVI(adata) + >>> vae.train() + >>> adata.obsm["X_scVI"] = vae.get_latent_representation() + >>> adata.obsm["X_normalized_scVI"] = vae.get_normalized_expression() + + Notes + ----- + See further usage examples in the following tutorials: + + 1. :doc:`/tutorials/notebooks/quick_start/api_overview` + 2. :doc:`/tutorials/notebooks/scrna/harmonization` + 3. :doc:`/tutorials/notebooks/scrna/scarches_scvi_tools` + 4. :doc:`/tutorials/notebooks/scrna/scvi_in_R` + + See Also + -------- + :class:`~scvi.module.VAE` + """ + + _module_cls = ASSAYVAE + _LATENT_QZM_KEY = "assayvi_latent_qzm" + _LATENT_QZV_KEY = "assayvi_latent_qzv" + _data_splitter_cls = DataSplitter + _training_plan_cls = AdversarialTrainingPlan + _train_runner_cls = TrainRunner + + def __init__( + self, + adata: AnnData | None = None, + n_hidden: int = 128, + n_latent: int = 10, + n_layers: int = 1, + dropout_rate: float = 0.05, + dispersion: Literal["gene", "gene-batch", "gene-cell"] = "gene", + gene_likelihood: Literal["zinb", "nb", "poisson", "normal"] = "nb", + prior: Literal["normal", "mog", "vamp"] = "normal", + pseudoinputs_data_indices: np.array | None = None, + n_prior_components: int = 50, + **kwargs, + ): + super().__init__(adata) + print('22222211') + + self._module_kwargs = { + "n_hidden": n_hidden, + "n_latent": n_latent, + "n_layers": n_layers, + "dropout_rate": dropout_rate, + "dispersion": dispersion, + "gene_likelihood": gene_likelihood, + **kwargs, + } + self._model_summary_string = ( + "SCVI model with the following parameters: \n" + f"n_hidden: {n_hidden}, n_latent: {n_latent}, n_layers: {n_layers}, " + f"dropout_rate: {dropout_rate}, dispersion: {dispersion}, " + f"gene_likelihood: {gene_likelihood}." + ) + + if prior == "vamp": + if pseudoinputs_data_indices is None: + pseudoinputs_data_indices = np.random.randint( + 0, self.summary_stats.n_cells, n_prior_components + ) + assert pseudoinputs_data_indices.shape[0] == n_prior_components + assert pseudoinputs_data_indices.ndim == 1 + pseudoinput_data = next( + iter( + self._make_data_loader( + adata=adata, + indices=pseudoinputs_data_indices, + batch_size=n_prior_components, + shuffle=False, + ) + ) + ) + else: + pseudoinput_data = None + + n_cats_per_cov = ( + self.adata_manager.get_state_registry(REGISTRY_KEYS.CAT_COVS_KEY).n_cats_per_key + if REGISTRY_KEYS.CAT_COVS_KEY in self.adata_manager.data_registry + else None + ) + self.module = self._module_cls( + n_input=self.summary_stats.n_vars, + n_batch=self.summary_stats.n_batch, + n_assay=self.summary_stats.n_assay, + n_labels=self.summary_stats.n_labels, + n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), + n_cats_per_cov=n_cats_per_cov, + n_hidden=n_hidden, + n_latent=n_latent, + n_layers=n_layers, + dropout_rate=dropout_rate, + dispersion=dispersion, + gene_likelihood=gene_likelihood, + prior=prior, + pseudoinput_data=pseudoinput_data, + n_prior_components=n_prior_components, + **kwargs, + ) + self.module.minified_data_type = self.minified_data_type + + self.init_params_ = self._get_init_params(locals()) + + @devices_dsp.dedent + def train( + self, + max_epochs: int | None = None, + lr: float = 4e-3, + accelerator: str = "auto", + devices: int | list[int] | str = "auto", + train_size: float | None = None, + validation_size: float | None = None, + shuffle_set_split: bool = True, + batch_size: int = 256, + early_stopping: bool = True, + check_val_every_n_epoch: int | None = None, + reduce_lr_on_plateau: bool = True, + n_steps_kl_warmup: int | None = None, + n_epochs_kl_warmup: int | None = None, + adversarial_classifier: bool | None = None, + adversarial_key: str = "assay", + datasplitter_kwargs: dict | None = None, + plan_kwargs: dict | None = None, + external_indexing: list[np.array] = None, + **kwargs, + ): + """Trains the model using amortized variational inference. + + Parameters + ---------- + max_epochs + Number of passes through the dataset. + lr + Learning rate for optimization. + %(param_accelerator)s + %(param_devices)s + train_size + Size of training set in the range [0.0, 1.0]. + validation_size + Size of the test set. If `None`, defaults to 1 - `train_size`. If + `train_size + validation_size < 1`, the remaining cells belong to a test set. + shuffle_set_split + Whether to shuffle indices before splitting. If `False`, the val, train, and test set + are split in the sequential order of the data according to `validation_size` and + `train_size` percentages. + batch_size + Minibatch size to use during training. + early_stopping + Whether to perform early stopping with respect to the validation set. + check_val_every_n_epoch + Check val every n train epochs. By default, val is not checked, unless `early_stopping` + is `True` or `reduce_lr_on_plateau` is `True`. If either of the latter conditions are + met, val is checked every epoch. + reduce_lr_on_plateau + Reduce learning rate on plateau of validation metric (default is ELBO). + n_epochs_kl_warmup + Number of epochs to scale weight on KL divergences from 0 to 1. + scale_adversarial_classifier + How to weight adversarial classifier in the latent space. This helps mixing when + there are multiple assays. Defaults to `1`. + adversarial_key + Key in `adata.obs` that corresponds to batch or assay key to use for adversarial + training. If `None`, defaults to the assay key. + datasplitter_kwargs + Additional keyword arguments passed into :class:`~scvi.dataloaders.DataSplitter`. + plan_kwargs + Keyword args for :class:`~scvi.train.AdversarialTrainingPlan`. Keyword arguments passed + to `train()` will overwrite values present in `plan_kwargs`, when appropriate. + external_indexing + A list of data split indices in the order of training, validation, and test sets. + Validation and test set are not required and can be left empty. + **kwargs + Other keyword args for :class:`~scvi.train.Trainer`. + """ + if adversarial_classifier is None: + if self.module.n_assay > 1: + adversarial_classifier = True + else: + adversarial_classifier = False + n_epochs_kl_warmup = ( + n_epochs_kl_warmup if n_epochs_kl_warmup is not None else max_epochs//2 + ) + if reduce_lr_on_plateau: + check_val_every_n_epoch = 1 + + update_dict = { + "lr": lr, + "adversarial_classifier": adversarial_classifier, + "adversarial_key": adversarial_key, + "reduce_lr_on_plateau": reduce_lr_on_plateau, + "n_epochs_kl_warmup": n_epochs_kl_warmup, + "n_steps_kl_warmup": n_steps_kl_warmup, + } + if plan_kwargs is not None: + plan_kwargs.update(update_dict) + else: + plan_kwargs = update_dict + + if max_epochs is None: + max_epochs = get_max_epochs_heuristic(self.adata.n_obs) + + plan_kwargs = plan_kwargs if isinstance(plan_kwargs, dict) else {} + datasplitter_kwargs = datasplitter_kwargs or {} + + data_splitter = self._data_splitter_cls( + self.adata_manager, + train_size=train_size, + validation_size=validation_size, + shuffle_set_split=shuffle_set_split, + batch_size=batch_size, + external_indexing=external_indexing, + **datasplitter_kwargs, + ) + training_plan = self._training_plan_cls(self.module, **plan_kwargs) + runner = self._train_runner_cls( + self, + training_plan=training_plan, + data_splitter=data_splitter, + max_epochs=max_epochs, + accelerator=accelerator, + devices=devices, + early_stopping=early_stopping, + check_val_every_n_epoch=check_val_every_n_epoch, + **kwargs, + ) + return runner() + + @classmethod + @setup_anndata_dsp.dedent + def setup_anndata( + cls, + adata: AnnData, + layer: str | None = None, + batch_key: str | None = None, + assay_key: str | None = None, + labels_key: str | None = None, + categorical_covariate_keys: list[str] | None = None, + continuous_covariate_keys: list[str] | None = None, + **kwargs, + ): + """%(summary)s. + + Parameters + ---------- + %(param_adata)s + %(param_layer)s + %(param_batch_key)s + assay_key + Key in ``adata.obs`` that corresponds to the assay of the data. + %(param_label_key)s + %(param_cat_cov_keys)s + %(param_cont_cov_keys)s + """ + setup_method_args = cls._get_setup_method_args(**locals()) + anndata_fields = [ + LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), + CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), + CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), + CategoricalObsField(REGISTRY_KEYS.LABELS_KEY, labels_key), + CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), + NumericalJointObsField(REGISTRY_KEYS.CONT_COVS_KEY, continuous_covariate_keys), + ] + # register new fields if the adata is minified + adata_minify_type = _get_adata_minify_type(adata) + if adata_minify_type is not None: + anndata_fields += cls._get_fields_for_adata_minification(adata_minify_type) + adata_manager = AnnDataManager(fields=anndata_fields, setup_method_args=setup_method_args) + adata_manager.register_fields(adata, **kwargs) + cls.register_manager(adata_manager) diff --git a/src/scvi/external/assayvi/_module.py b/src/scvi/external/assayvi/_module.py new file mode 100644 index 0000000000..391c2026e5 --- /dev/null +++ b/src/scvi/external/assayvi/_module.py @@ -0,0 +1,712 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +import torch +from torch import distributions +from torch.nn.functional import one_hot + +from scvi import REGISTRY_KEYS +from scvi.data._constants import ADATA_MINIFY_TYPE +from scvi.module._classifier import Classifier +from scvi.module._constants import MODULE_KEYS +from scvi.module.base import ( + BaseMinifiedModeModuleClass, + EmbeddingModuleMixin, + LossOutput, + MogPrior, + StandardPrior, + VampPrior, + 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 ASSAYVAE(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-assay"``: parameter is constant per gene per assay. + * ``"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`. + 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 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`). + 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"``. + + Notes + ----- + Lifecycle: argument ``batch_representation`` is experimental in v1.2. + """ + + def __init__( + self, + n_input: int, + n_batch: int = 0, + n_assay: 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-assay", "gene-cell"] = "gene", + log_variational: bool = True, + gene_likelihood: Literal["zinb", "nb", "poisson"] = "nb", + latent_distribution: Literal["normal", "ln"] = "normal", + encode_covariates: bool = False, + encode_assay: bool = True, + deeply_inject_covariates: bool = False, + batch_representation: Literal["one-hot", "embedding"] = "one-hot", + use_batch_norm: Literal["encoder", "decoder", "none", "both"] = "none", + use_layer_norm: Literal["encoder", "decoder", "none", "both"] = "both", + 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, + conditional_norm: bool = True, + conditional_output: bool = True, + prior: str | None = None, + pseudoinput_data: dict | None = None, + n_prior_components: int | None = 30, + mmd_kernel: str = "rbf", + ): + 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_assay = n_assay + self.n_labels = n_labels + self.latent_distribution = latent_distribution + self.encode_covariates = encode_covariates + self.use_observed_lib_size = True + + + if self.dispersion == "gene": + self.px_r = torch.nn.Parameter(3.*torch.ones(n_input)) + elif self.dispersion == "gene-batch": + self.px_r = torch.nn.Parameter(3.*torch.ones(n_input, n_batch)) + elif self.dispersion == "gene-assay": + self.px_r = torch.nn.Parameter(3.*torch.ones(n_input, n_assay)) + elif self.dispersion == "gene-cell": + pass + else: + raise ValueError( + "`dispersion` must be one of 'gene', 'gene-batch', 'gene-assay', 'gene-cell'." + ) + + self.batch_representation = batch_representation + n_cats_per_cov_ = list([] if n_cats_per_cov is None else n_cats_per_cov) + n_continuous = n_continuous_cov + + if self.batch_representation == "embedding": + self.init_embedding(REGISTRY_KEYS.BATCH_KEY, n_batch, **(batch_embedding_kwargs or {})) + n_continuous += self.get_embedding_dim(REGISTRY_KEYS.BATCH_KEY) + 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" + + if self.batch_representation == "embedding": + cat_list = n_cats_per_cov_ + else: + cat_list = [n_batch] + n_cats_per_cov_ + self.encode_assay = encode_assay + self.batch_representation_encoder = False + conditional_category = 0 + if encode_assay: + encode_assay_list = [n_assay] + else: + encode_assay_list = [0] + if not encode_covariates: + encoder_cat_list = encode_assay_list + n_cont_encoder = 0 + else: + if not conditional_norm and self.batch_representation == "embedding": + # If we are not using conditional norm, we need to pass the batch information + self.batch_representation_encoder = True + encoder_cat_list = encode_assay_list + n_cats_per_cov_ + n_cont_encoder = n_continuous + else: + encoder_cat_list = encode_assay_list + [n_batch] + n_cats_per_cov_ + n_cont_encoder = n_continuous_cov + if conditional_norm and not encode_assay: + conditional_category = 1 + _extra_encoder_kwargs = extra_encoder_kwargs or {} + self.z_encoder = Encoder( + n_input, + n_latent, + n_continuous=n_cont_encoder, + 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, + conditional_norm=conditional_norm, + conditional_category=conditional_category, + **_extra_encoder_kwargs, + ) + + _extra_decoder_kwargs = extra_decoder_kwargs or {} + self.decoder = DecoderSCVI( + n_latent, + n_input, + n_cat_list=cat_list, + n_continuous=n_continuous, + n_layers=n_layers, + n_hidden=n_hidden, + n_conditions_output=self.n_assay if conditional_output else 0, + inject_covariates=deeply_inject_covariates, + use_batch_norm=use_batch_norm_decoder, + use_layer_norm=use_layer_norm_decoder, + scale_activation="softmax", + **_extra_decoder_kwargs, + ) + + cls_parameters = { + "n_layers": 1, + "n_hidden": 128, + "dropout_rate": 0.0, + "logits": True, + } + + if self.n_labels > 1: + self.classifier = Classifier( + n_latent, + n_labels=self.n_labels, + use_batch_norm=False, + use_layer_norm=True, + **cls_parameters, + ) + if prior == "gaussian": + self.prior = StandardPrior() + elif prior == "vamp": + assert pseudoinput_data is not None, ( + "Pseudoinput data must be specified if using VampPrior" + ) + pseudoinput_data = self._get_inference_input( + pseudoinput_data, + full_forward_pass=True + ) + print('include training.') + cat_list = [n_batch] + n_cats_per_cov_ + encode_assay_list + self.prior = VampPrior( + n_components=n_prior_components, + inference=self._regular_inference, + encoder=self.z_encoder, + pseudoinputs=pseudoinput_data, + n_cat_list=cat_list, + trainable_priors=True, + additional_categorical_covariates=["assay_index"] + ) + elif prior == "mog": + self.prior = MogPrior( + n_components=n_prior_components, + n_latent=n_latent, + ) + elif prior == "mog_celltype": + self.prior = MogPrior( + n_components=n_labels, + n_latent=n_latent, + celltype_bias=True + ) + else: + raise ValueError( + "`prior` must be one of 'gaussian', 'vamp', 'mog', 'mog_celltype'.") + + 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.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), + 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.""" + 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.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), + 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), + } + + @auto_move_data + def _regular_inference( + self, + x: torch.Tensor, + batch_index: torch.Tensor, + assay_index: torch.Tensor | None = None, + 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_ = x_/x_.mean(1).unsqueeze(1) + x_ = torch.log1p(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.encode_covariates and self.batch_representation_encoder: + batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index) + if cont_covs is not None: + cont_input = torch.cat([cont_covs, batch_rep], dim=-1) + else: + cont_input = batch_rep + else: + cont_input = cont_covs + if not self.encode_assay: + assay_index = None + else: + assay_index = assay_index.long() + qz, z = self.z_encoder(x_, assay_index, batch_index, *categorical_input, cont_input=cont_input) + + if n_samples > 1: + untran_z = qz.sample((n_samples,)) + z = self.z_encoder.z_transformation(untran_z) + 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.LIBRARY_KEY: library, + } + + @auto_move_data + def generative( + self, + z: torch.Tensor, + library: torch.Tensor, + batch_index: torch.Tensor, + assay_index: torch.Tensor | None = None, + cont_covs: torch.Tensor | None = None, + cat_covs: torch.Tensor | None = None, + size_factor: torch.Tensor | None = None, # Consistency + y: torch.Tensor | None = None, + transform_batch: torch.Tensor | None = None, + transform_assay: 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, + ) + + 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 transform_assay is not None: + assay_index = torch.ones_like(assay_index) * transform_assay + + if self.batch_representation == "embedding": + batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index) + if cont_covs is not None: + cont_input = torch.cat([cont_covs, batch_rep], dim=-1) + else: + cont_input = batch_rep + else: + cont_input = cont_covs + + px_scale, px_r, px_rate, px_dropout = self.decoder( + self.dispersion, + z, + library, + batch_index, + *categorical_input, + y, + cont_input=cont_input, + output_condition=assay_index.long(), + ) + + if self.dispersion == "gene-assay": + px_r = linear( + one_hot(assay_index.squeeze(-1).long(), self.n_assay).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: float = 1.0, + weight_assay_loss: float = 0.0, + weight_global: float = 1.0, + weight_kl_sample: float = 1.0, + classification_ratio: float = 500., + ) -> LossOutput: + """Compute the loss.""" + from torch.distributions import kl_divergence + + x = tensors[REGISTRY_KEYS.X_KEY] + y = tensors[REGISTRY_KEYS.LABELS_KEY] + + kl_divergence_z = self.prior.kl( + qz=inference_outputs[MODULE_KEYS.QZ_KEY], + z=inference_outputs[MODULE_KEYS.Z_KEY], + labels=y, + ) + + if self.get_embedding_variational(REGISTRY_KEYS.BATCH_KEY, default_value=False): + pz_sample = self.compute_embedding( + REGISTRY_KEYS.BATCH_KEY, tensors[REGISTRY_KEYS.BATCH_KEY], return_dist=True) + qz_sample = distributions.Normal( + torch.zeros_like(pz_sample.loc), torch.ones_like(pz_sample.scale)) + kl_divergence_sample = kl_divergence(qz_sample, pz_sample).sum(dim=1) + else: + kl_divergence_sample = torch.zeros_like(kl_divergence_z) + + assay_index = tensors.get(REGISTRY_KEYS.ASSAY_KEY, None) + if weight_assay_loss > 0.0 and assay_index is not None: + assay_loss = self._compute_assay_penalty( + inference_outputs[MODULE_KEYS.QZ_KEY].loc, + assay_index + ) + else: + assay_loss = 0.0 + + reconst_loss = -generative_outputs[MODULE_KEYS.PX_KEY].log_prob(x).sum(-1) + + kl_global = torch.zeros_like(kl_divergence_z) + + if self.gene_likelihood == "zinb": + zi = generative_outputs[MODULE_KEYS.PX_KEY].zi_logits + kl_global -= distributions.Exponential( + 10.*torch.ones_like(zi)).log_prob(torch.exp(zi)).sum(-1) + if self.gene_likelihood == "zinb" or self.gene_likelihood == "nb": + theta = generative_outputs[MODULE_KEYS.PX_KEY].theta + kl_global -= distributions.Exponential( + torch.ones_like(theta)).log_prob(1/theta).sum(-1) + + weighted_kl_local = kl_weight * (kl_divergence_z + weight_kl_sample * kl_divergence_sample) + + loss = torch.mean( + reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss + + weight_global * kl_global) + + if self.n_labels > 1: + logits = self.classifier(inference_outputs[MODULE_KEYS.Z_KEY]) + classification_loss_ = torch.nn.functional.cross_entropy(logits, y.ravel(), reduction="none") + mask = (y != self.n_labels) + classification_loss = classification_ratio * torch.masked_select( + classification_loss_, mask).mean(0) + loss += torch.mean(classification_loss) + + return LossOutput( + loss=loss, reconstruction_loss=reconst_loss, kl_local=kl_divergence_z, + classification_loss=classification_loss, logits=logits, true_labels=y + ) + + return LossOutput( + loss=loss, + reconstruction_loss=reconst_loss, + kl_local={ + MODULE_KEYS.KL_Z_KEY: kl_divergence_z, + MODULE_KEYS.KL_SAMPLE_KEY: kl_divergence_sample, + }, + ) + + @torch.inference_mode() + def sample( + self, + tensors: dict[str, torch.Tensor], + n_samples: int = 1, + max_poisson_rate: float = 1e8, + ) -> 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 :meth:`~scvi.module.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. + + 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, compute_loss=False + ) + + dist = generative_outputs[MODULE_KEYS.PX_KEY] + if self.gene_likelihood == "poisson": + dist = Poisson(torch.clamp(dist.rate, max=max_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() + + def _compute_assay_penalty( + self, params, assay): + assay = assay.squeeze(-1).long() + unique = torch.unique(assay) + pair_penalty = torch.tensor(0., device=assay.device) + if len(unique) > 1: + for i in unique: + pp = self.mmd(params, mask=(assay == i)) + pair_penalty += pp + + return pair_penalty + + def mmd(self, params, mask=None): + if mask is not None: + mod_1 = params[mask] + mod_2 = params[~mask] + return rbf_kernel(mod_1, mod_2) + + +@auto_move_data +def rbf_kernel(x, y, gammas=None): + """ + Compute the RBF kernel between two tensors. + + Parameters + ---------- + x : torch.Tensor + Input tensor of shape (N, D). + y : torch.Tensor + Input tensor of shape (M, D). + gammas : list of float or None + List of gamma values to compute the kernel with. + + Returns + ------- + kernel_sum : torch.Tensor + Tensor of shape (N, M), the sum of RBF kernels over all gamma values. + """ + if gammas is None: + gammas = [ + 1e-10, + 1e-8, + 1e-6, + 1e-4, + 1e-3, + 1e-2, + 1e-1, + 1, + 2, + 5, + 10, + ] + kxy = torch.cdist(x, y).pow(2) + kxx = torch.cdist(x, x).pow(2) + kyy = torch.cdist(y, y).pow(2) + + kernel_sum_xy = torch.tensor(0.0, device=x.device) + kernel_sum_xx = torch.tensor(0.0, device=x.device) + kernel_sum_yy = torch.tensor(0.0, device=x.device) + for gamma in gammas: + kernel_sum_xy += torch.exp(-gamma * kxy).mean() + kernel_sum_xx += torch.exp(-gamma * kxx).mean() + kernel_sum_yy += torch.exp(-gamma * kyy).mean() + + return (kernel_sum_xx + kernel_sum_yy - 2 * kernel_sum_xy) / len(gammas) diff --git a/src/scvi/model/base/_embedding_mixin.py b/src/scvi/model/base/_embedding_mixin.py index 7849be92ac..03a0626501 100644 --- a/src/scvi/model/base/_embedding_mixin.py +++ b/src/scvi/model/base/_embedding_mixin.py @@ -13,7 +13,7 @@ class EmbeddingMixin: - """``EXPERIMENTAL`` Mixin class for initializing and using embeddings in a model. + """Mixin class for computing covariate embeddings of a model. Must be used with a module that inherits from :class:`~scvi.module.base.EmbeddingModuleMixin`. @@ -25,14 +25,34 @@ def get_batch_representation( adata: AnnData | None = None, indices: list[int] | None = None, batch_size: int | None = None, + key: str = REGISTRY_KEYS.BATCH_KEY, + return_mean: bool = True, ) -> np.ndarray: - """Get the batch representation for a given set of indices.""" + """Get the batch representation for a given set of indices. + + Parameters + ---------- + adata + AnnData object to use. + indices + Indices to get the batch representation for. + batch_size + Minibatch size for computing the batch representation. + key + Setup key to compute the batch representation for. + return_mean + Return the mean of the batch representation. Or sample from it. + """ if not isinstance(self.module, EmbeddingModuleMixin): raise ValueError("The current `module` must inherit from `EmbeddingModuleMixin`.") + if key not in self.module.embeddings_dim: + raise ValueError(f"Embedding {key} not found. Enable it during model setup.") adata = self._validate_anndata(adata) dataloader = self._make_data_loader(adata=adata, indices=indices, batch_size=batch_size) - key = REGISTRY_KEYS.BATCH_KEY - tensors = [self.module.compute_embedding(key, tensors[key]) for tensors in dataloader] + tensors = [ + self.module.compute_embedding(key, tensors[key], return_mean=return_mean) + for tensors in dataloader + ] return torch.cat(tensors).detach().cpu().numpy() diff --git a/src/scvi/module/_constants.py b/src/scvi/module/_constants.py index f4fabc4ac1..4a75333db0 100644 --- a/src/scvi/module/_constants.py +++ b/src/scvi/module/_constants.py @@ -24,6 +24,7 @@ class _MODULE_KEYS(NamedTuple): # loss KL_L_KEY: str = "kl_divergence_l" KL_Z_KEY: str = "kl_divergence_z" + KL_SAMPLE_KEY: str = "kl_divergence_sample" MODULE_KEYS = _MODULE_KEYS() diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index e9277d0986..1eda3831ef 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -148,7 +148,6 @@ def __init__( self, n_input: int, n_batch: int = 0, - n_assay: int = 0, n_labels: int = 0, n_hidden: int = 128, n_latent: int = 10, @@ -188,8 +187,6 @@ def __init__( self.log_variational = log_variational self.gene_likelihood = gene_likelihood self.n_batch = n_batch - self.n_input = n_input - self.n_assay = n_assay self.n_labels = n_labels self.n_hidden = n_hidden self.n_layers = n_layers @@ -241,13 +238,8 @@ def __init__( 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) - if n_assay > 1: - encoder_cat_list = [n_assay] + list([] if n_cats_per_cov is None else n_cats_per_cov) - n_input_encoder = n_input + n_continuous_cov * encode_covariates - else: - encoder_cat_list = cat_list - encoder_cat_list = encoder_cat_list if encode_covariates else None + encoder_cat_list = cat_list if encode_covariates else None _extra_encoder_kwargs = extra_encoder_kwargs or {} self.z_encoder = Encoder( n_input_encoder, @@ -291,7 +283,6 @@ def __init__( n_cat_list=cat_list, n_layers=n_layers, n_hidden=n_hidden, - n_assay=self.n_assay, inject_covariates=deeply_inject_covariates, use_batch_norm=use_batch_norm_decoder, use_layer_norm=use_layer_norm_decoder, @@ -422,7 +413,7 @@ def _regular_inference( elif self.encode_covariates and self.batch_representation == "embedding": 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, batch_index, *categorical_input) + qz, z = self.z_encoder(encoder_input, *categorical_input) else: qz, z = self.z_encoder(encoder_input, batch_index, *categorical_input) diff --git a/src/scvi/module/base/_embedding_mixin.py b/src/scvi/module/base/_embedding_mixin.py index 1286d27aef..7651358023 100644 --- a/src/scvi/module/base/_embedding_mixin.py +++ b/src/scvi/module/base/_embedding_mixin.py @@ -1,4 +1,5 @@ import torch +from torch.distributions import Normal from torch.nn import ModuleDict from scvi.module.base._decorators import auto_move_data @@ -6,7 +7,7 @@ class EmbeddingModuleMixin: - """``EXPERIMENTAL`` Mixin class for initializing and using embeddings in a module.""" + """Mixin class for initializing and using embeddings in a module.""" @property def embeddings_dict(self) -> ModuleDict: @@ -15,10 +16,25 @@ def embeddings_dict(self) -> ModuleDict: self._embeddings_dict = ModuleDict() return self._embeddings_dict + @property + def embeddings_dim(self) -> dict: + """Dictionary of embeddings dimensions.""" + if not hasattr(self, "_embeddings_dict"): + self._embeddings_dim = {} + return self._embeddings_dim + + @property + def variational(self) -> dict: + """Dictionary of whether embedding is variational.""" + if not hasattr(self, "_variational"): + self._variational = {} + return self._variational + def add_embedding(self, key: str, embedding: Embedding, overwrite: bool = False) -> None: """Add an embedding to the module.""" if key in self.embeddings_dict and not overwrite: raise KeyError(f"Embedding {key} already exists.") + torch.nn.init.zeros_(embedding.weight) self.embeddings_dict[key] = embedding def remove_embedding(self, key: str) -> None: @@ -27,24 +43,68 @@ def remove_embedding(self, key: str) -> None: raise KeyError(f"Embedding {key} not found.") del self.embeddings_dict[key] - def get_embedding(self, key: str) -> Embedding: + def get_embedding(self, key: str,) -> Embedding: """Get an embedding from the module.""" if key not in self.embeddings_dict: raise KeyError(f"Embedding {key} not found.") return self.embeddings_dict[key] + def get_embedding_dim(self, key: str, default_value: str | None = None) -> int: + """Get the dimension of an embedding.""" + if key not in self.embeddings_dim: + if default_value is not None: + return default_value + else: + raise KeyError(f"Embedding {key} not found.") + return self.embeddings_dim[key] + + def get_embedding_variational(self, key: str, default_value: str | None = None) -> bool: + """Get whether an embedding is variational.""" + if key not in self.variational: + if default_value is not None: + return default_value + else: + raise KeyError(f"Embedding {key} not found.") + return self.variational[key] + def init_embedding( self, key: str, num_embeddings: int, embedding_dim: int = 5, + variational: bool = False, **kwargs, ) -> None: """Initialize an embedding in the module.""" + self.embeddings_dim[key] = embedding_dim + self.variational[key] = variational + + if variational: + embedding_dim *= 2 self.add_embedding(key, Embedding(num_embeddings, embedding_dim, **kwargs)) @auto_move_data - def compute_embedding(self, key: str, indices: torch.Tensor) -> torch.Tensor: + def compute_embedding( + self, + key: str, + indices: torch.Tensor, + return_mean: bool = False, + return_dist: bool = False, + ) -> torch.Tensor: """Forward pass for an embedding.""" indices = indices.flatten() if indices.ndim > 1 else indices - return self.get_embedding(key)(indices) + embedding = self.get_embedding(key)(indices) + if self.get_embedding_variational(key): + embedding_dim = self.get_embedding_dim(key) + if return_mean: + return embedding[:, :embedding_dim] + dist = Normal( + embedding[:, :embedding_dim], + torch.exp(embedding[:, embedding_dim:]) + ) + if return_dist: + return dist + else: + return dist.rsample() + else: + return embedding diff --git a/src/scvi/module/base/_priors.py b/src/scvi/module/base/_priors.py index 6778c908a5..5fd0dcab8d 100644 --- a/src/scvi/module/base/_priors.py +++ b/src/scvi/module/base/_priors.py @@ -96,7 +96,7 @@ def __init__( ): super().__init__() self.prior_means = torch.nn.Parameter(0.1 * torch.randn([n_components, n_latent])) - self.prior_log_scales = torch.nn.Parameter(torch.zeros([n_components, n_latent]) - 1.0) + self.prior_log_scales = torch.nn.Parameter(torch.zeros([n_components, n_latent]) - 1.) self.prior_logits = torch.nn.Parameter(torch.zeros([n_components])) self.celltype_bias = celltype_bias if celltype_bias: diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index b328c26b6f..6611790723 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -55,6 +55,9 @@ class FCLayers(nn.Module): The dimensionality of the input n_out The dimensionality of the output + n_continuous + The dimensionality of the continuous covariates + including batch embeddings. n_cat_list A list containing, for each category of interest, the number of categories. Each category will be @@ -85,6 +88,7 @@ def __init__( self, n_in: int, n_out: int, + n_continuous: int = 0, n_cat_list: Iterable[int] = None, n_cont: int = 0, n_layers: int = 1, @@ -97,6 +101,7 @@ def __init__( inject_covariates: bool = True, activation_fn: nn.Module = nn.ReLU, conditional_norm: bool = False, + conditional_category: int = 0 ): super().__init__() self.inject_covariates = inject_covariates @@ -107,6 +112,9 @@ def __init__( self.n_cat_list = [n_cat if n_cat > 1 else 0 for n_cat in n_cat_list] else: self.n_cat_list = [] + self.n_continuous = n_continuous + + self.cond_cat = conditional_category self.n_cov = n_cont + sum(self.n_cat_list) @@ -117,18 +125,18 @@ def __init__( f"Layer {i}", nn.Sequential( nn.Linear( - n_in + self.n_cov * self.inject_into_layer(i), + n_in + (cat_dim + n_continuous) * self.inject_into_layer(i), n_out, bias=bias, ), - # non-default params come from defaults in the original Tensorflow - # implementation - ConditionalBatchNorm2d(n_out, self.n_cat_list[0], momentum=0.01, eps=0.001) + # non-default params come from defaults in Tensorflow implementation + ConditionalBatchNorm2d( + n_out, self.n_cat_list[self.cond_cat], momentum=0.01, eps=0.001) if conditional_norm and use_batch_norm else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) if use_batch_norm else None, - # non-default params come from defaults in original Tensorflow implementation - ConditionalLayerNorm(n_out, self.n_cat_list[0]) + # non-default params come from defaults in Tensorflow implementation + ConditionalLayerNorm(n_out, self.n_cat_list[self.cond_cat]) if conditional_norm and use_layer_norm else nn.LayerNorm(n_out, elementwise_affine=False) if use_layer_norm else None, @@ -175,7 +183,7 @@ def _hook_fn_zero_out(grad): b = layer.bias.register_hook(_hook_fn_zero_out) self.hooks.append(b) - def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = None): + def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | None = None): """Forward computation on ``x``. Parameters @@ -184,8 +192,8 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = N tensor of values with shape ``(n_in,)`` cat_list list of category membership(s) for this sample - cont - tensor of continuous covariates with shape ``(n_cont,)`` + cont_input + tensor of continuous covariates with shape ``(n_continuous,)`` Returns ------- @@ -198,6 +206,8 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = N if len(self.n_cat_list) > len(cat_list): raise ValueError("nb. categorical args provided doesn't match init. params.") + if self.n_continuous>0 and cont_input.shape[-1] != self.n_continuous: + raise ValueError("continuous dims provided doesn't match init. params.") for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): if n_cat and cat is None: raise ValueError("cat not provided while n_cat != 0 in init. params.") @@ -207,7 +217,8 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = N else: one_hot_cat = cat # cat has already been one_hot encoded one_hot_cat_list += [one_hot_cat] - cov_list = cont_list + one_hot_cat_list + if cont_input is not None: + one_hot_cat_list += [cont_input] for i, layers in enumerate(self.fc_layers): for layer in layers: if layer is not None: @@ -215,10 +226,12 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = N layer, ConditionalLayerNorm): if x.dim() == 3: x = torch.cat( - [(layer(slice_x, cat_list[0])).unsqueeze(0) for slice_x in x], dim=0 + [(layer(slice_x, cat_list[self.cond_cat])).unsqueeze(0) + for slice_x in x], + dim=0 ) else: - x = layer(x=x, y=cat_list[0]) + x = layer(x=x, y=cat_list[self.cond_cat]) elif isinstance(layer, nn.BatchNorm1d): if x.dim() == 3: if ( @@ -259,7 +272,10 @@ class Encoder(nn.Module): The dimensionality of the input (data space) n_output The dimensionality of the output (latent space) - n_cat_list + n_continuous + The dimensionality of the continuous covariates + including batch embeddings. + n_cat_l)t A list containing the number of categories for each category of interest. Each category will be included using a one-hot encoding @@ -287,7 +303,8 @@ def __init__( self, n_input: int, n_output: int, - n_cat_list: Iterable[int] = None, + n_continuous: int = 0, + n_cat_list: Iterable[int] | None = None, n_layers: int = 1, n_hidden: int = 128, dropout_rate: float = 0.1, @@ -304,6 +321,7 @@ def __init__( self.encoder = FCLayers( n_in=n_input, n_out=n_hidden, + n_continuous=n_continuous, n_cat_list=n_cat_list, n_layers=n_layers, n_hidden=n_hidden, @@ -320,7 +338,12 @@ def __init__( self.z_transformation = _identity self.var_activation = torch.exp if var_activation is None else var_activation - def forward(self, x: torch.Tensor, *cat_list: int): + def forward( + self, + x: torch.Tensor, + *cat_list: int, + cont_input: torch.Tensor | None = None, + ): r"""The forward computation for a single sample. #. Encodes the data into latent space using the encoder network @@ -334,6 +357,8 @@ def forward(self, x: torch.Tensor, *cat_list: int): tensor with shape (n_input,) cat_list list of category membership(s) for this sample + cont_input + optional tensor with shape (n_continuous,) Returns ------- @@ -342,7 +367,7 @@ def forward(self, x: torch.Tensor, *cat_list: int): """ # Parameters for latent distribution - q = self.encoder(x, *cat_list) + q = self.encoder(x, *cat_list, cont_input=cont_input) q_m = self.mean_encoder(q) q_v = self.var_activation(self.var_encoder(q)) + self.var_eps dist = Normal(q_m, q_v.sqrt()) @@ -364,6 +389,8 @@ class DecoderSCVI(nn.Module): The dimensionality of the input (latent space) n_output The dimensionality of the output (data space) + n_continuous + The dimensionality of the continuous covariates n_cat_list A list containing the number of categories for each category of interest. Each category will be @@ -372,6 +399,8 @@ class DecoderSCVI(nn.Module): The number of fully-connected hidden layers n_hidden The number of nodes per hidden layer + n_conditions_output + The number of conditions add to the scale and dropout parameters. dropout_rate Dropout rate to apply to each of the hidden layers inject_covariates @@ -390,10 +419,11 @@ def __init__( self, n_input: int, n_output: int, + n_continuous: int = 0, n_cat_list: Iterable[int] = None, n_layers: int = 1, n_hidden: int = 128, - n_assay: int = 1, + n_conditions_output: int = 0, inject_covariates: bool = True, use_batch_norm: bool = False, use_layer_norm: bool = False, @@ -401,9 +431,11 @@ def __init__( **kwargs, ): super().__init__() + self.n_conditions_output = n_conditions_output self.px_decoder = FCLayers( n_in=n_input, n_out=n_hidden, + n_continuous=n_continuous, n_cat_list=n_cat_list, n_layers=n_layers, n_hidden=n_hidden, @@ -421,18 +453,16 @@ def __init__( px_scale_activation = nn.Softplus() elif scale_activation == "exp": px_scale_activation = ExpActivation() + + # scale self.px_scale_decoder = nn.Sequential( - nn.Linear(n_hidden + n_assay, n_output), + nn.Linear(n_hidden + n_conditions_output, n_output), px_scale_activation, ) - # dispersion: here we only deal with gene-cell dispersion case - self.px_r_decoder = nn.Linear(n_hidden + n_assay, n_output) - + self.px_r_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) # dropout - self.px_dropout_decoder = nn.Linear(n_hidden + n_assay, n_output) - - self.n_assay = n_assay + self.px_dropout_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) def forward( self, @@ -440,7 +470,8 @@ def forward( z: torch.Tensor, library: torch.Tensor, *cat_list: int, - assay: torch.Tensor | None = None, + cont_input: torch.Tensor | None = None, + output_condition: torch.Tensor | None = None, ): """The forward computation for a single sample. @@ -457,14 +488,16 @@ def forward( * ``'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 - assay - tensor with shape ``(n_input,)`` of assay column z tensor with shape ``(n_input,)`` library library size cat_list list of category membership(s) for this sample + cont_input + tensor with shape ``(n_continuous,)`` + output_condition + tensor with shape ``(n_input,)`` used for conditioning the output layer Returns ------- @@ -473,11 +506,14 @@ def forward( """ # The decoder returns values for the parameters of the ZINB distribution - px = self.px_decoder(z, *cat_list) - if assay is not None: - one_hot_cat = nn.functional.one_hot(assay.squeeze(-1), self.n_assay) + px = self.px_decoder(z, *cat_list, cont_input=cont_input) + if output_condition is not None and self.n_conditions_output: + one_hot_cat = nn.functional.one_hot( + output_condition.squeeze(-1), + self.n_conditions_output + ) else: - one_hot_cat = torch.zeros(px.size(0), self.n_assay) + one_hot_cat = torch.zeros(px.size(0), self.n_conditions_output).to(px.device) px_cat = torch.cat([px, one_hot_cat], dim=-1) px_scale = self.px_scale_decoder(px_cat) px_dropout = self.px_dropout_decoder(px_cat) @@ -625,7 +661,7 @@ class MultiEncoder(nn.Module): def __init__( self, n_heads: int, - n_input_list: list[int], + n_input_list: Iterable[int], n_output: int, n_hidden: int = 128, n_layers_individual: int = 1, diff --git a/src/scvi/nn/_embedding.py b/src/scvi/nn/_embedding.py index 58a2740e1c..956109de2b 100644 --- a/src/scvi/nn/_embedding.py +++ b/src/scvi/nn/_embedding.py @@ -19,15 +19,15 @@ def _partial_freeze_hook_factory(freeze: int) -> Callable[[torch.Tensor], torch. """ def _partial_freeze_hook(grad: torch.Tensor) -> torch.Tensor: - grad = grad.clone() - grad[:freeze] = 0.0 - return grad + grad_copy = grad.clone() + grad_copy[:freeze] = 0.0 + return grad_copy return _partial_freeze_hook class Embedding(nn.Embedding): - """``EXPERIMENTAL`` Embedding layer with utility methods for extending.""" + """Embedding layer with utility methods for extending.""" @classmethod def extend( diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index aef615bd4b..59ee13517a 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -12,13 +12,17 @@ import torchmetrics.functional as tmf from lightning.pytorch.strategies.ddp import DDPStrategy from pyro.nn import PyroModule +from torch.distributions import Normal from torch.optim.lr_scheduler import ReduceLROnPlateau from scvi import REGISTRY_KEYS, settings from scvi.module import Classifier +from scvi.module._constants import MODULE_KEYS from scvi.module.base import ( BaseModuleClass, LossOutput, + MogPrior, + VampPrior, PyroBaseModuleClass, ) from scvi.train._constants import METRIC_KEYS @@ -573,6 +577,8 @@ class AdversarialTrainingPlan(TrainingPlan): Minimum learning rate allowed adversarial_classifier Whether to use adversarial classifier in the latent space + adversarial_key + Key in setup args to use for adversarial training. scale_adversarial_loss Scaling factor on the adversarial components of the loss. By default, adversarial loss is scaled from 1 to 0 following the opposite of @@ -603,6 +609,8 @@ def __init__( ] = "elbo_validation", lr_min: float = 0, adversarial_classifier: bool | Classifier = False, + adversarial_key: str = "batch", + adversarial_steps: int = 1, scale_adversarial_loss: float | Literal["auto"] = "auto", compile: bool = False, compile_kwargs: dict | None = None, @@ -626,26 +634,33 @@ def __init__( compile_kwargs=compile_kwargs, **loss_kwargs, ) + self.adversarial_key = adversarial_key if adversarial_classifier is True: - if self.module.n_batch == 1: + self.adversarial_steps = adversarial_steps + self.n_adversarial_name = f"n_{adversarial_key}" + if hasattr(self.module, self.n_adversarial_name): + self.n_output_classifier = getattr(self.module, self.n_adversarial_name) + else: + raise ValueError( + f"Adversarial key {adversarial_key} not found in module setup args." + ) + if self.n_output_classifier == 1: warnings.warn( - "Disabling adversarial classifier.", + "Disabling adversarial classifier as there is only one class.", UserWarning, stacklevel=settings.warnings_stacklevel, ) self.adversarial_classifier = False else: - self.n_output_classifier = self.module.n_assay self.adversarial_classifier = Classifier( n_input=self.module.n_latent, n_hidden=128, n_labels=self.n_output_classifier, - n_layers=1, + n_layers=2, logits=True, use_batch_norm=False, use_layer_norm=True, ) - print('new classifier') else: self.adversarial_classifier = adversarial_classifier self.scale_adversarial_loss = scale_adversarial_loss @@ -654,18 +669,17 @@ def __init__( def loss_adversarial_classifier(self, z, batch_index, predict_true_class=True): """Loss for adversarial classifier.""" n_classes = self.n_output_classifier - cls_logits = torch.nn.LogSoftmax(dim=1)(self.adversarial_classifier(z)) + cls_logits = self.adversarial_classifier(z) if predict_true_class: - cls_target = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes) + cls_target = batch_index.squeeze(-1) + loss = torch.nn.functional.cross_entropy(cls_logits, cls_target) else: - one_hot_batch = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes) - # place zeroes where the true label is - cls_target = (~one_hot_batch.bool()).float() - cls_target = cls_target / (n_classes - 1) - - l_soft = cls_logits * cls_target - loss = -l_soft.sum(dim=1).mean() + one_hot_batch = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes).float() + cls_target = (1 - one_hot_batch) / (n_classes - 1) + loss = - ( + cls_target * torch.nn.functional.log_softmax(cls_logits, dim=1) + ).sum(dim=1).mean() return loss @@ -679,7 +693,7 @@ def training_step(self, batch, batch_idx): if self.scale_adversarial_loss == "auto" else self.scale_adversarial_loss ) - batch_tensor = batch[REGISTRY_KEYS.ASSAY_KEY].long() + batch_tensor = batch[self.adversarial_key].long() opts = self.optimizers() if not isinstance(opts, list): @@ -711,11 +725,20 @@ def training_step(self, batch, batch_idx): # train adversarial classifier # this condition will not be met unless self.adversarial_classifier is not False if opt2 is not None: - loss = self.loss_adversarial_classifier(z.detach(), batch_tensor, True) - loss *= kappa - opt2.zero_grad() - self.manual_backward(loss) - opt2.step() + for _ in range(self.adversarial_steps): + qz = inference_outputs["qz"] + z = qz.sample().detach() + loss = self.loss_adversarial_classifier(z, batch_tensor, True) + if isinstance(self.module.prior, MogPrior) or isinstance(self.module.prior, VampPrior): + qz_m, qz_v = qz.loc.detach(), qz.scale.detach() + loss += self.module.prior.kl( + qz=Normal(qz_m, qz_v), + z=z, + labels=batch[REGISTRY_KEYS.LABELS_KEY].long(), + ).mean() + opt2.zero_grad() + self.manual_backward(loss) + opt2.step() # next part is for the usage of scib-metrics autotune with scvi if scvi_loss.extra_metrics is not None and len(scvi_loss.extra_metrics.keys()) > 0: From 3f7ef32201c4ade9e07d187afbb9636cdc343ec4 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 28 May 2025 16:04:47 -0700 Subject: [PATCH 04/24] Undo changes in scvi --- src/scvi/external/assayvi/_model.py | 13 +- src/scvi/external/assayvi/_module.py | 2 +- src/scvi/model/_scvi.py | 8 +- src/scvi/model/base/_training_mixin.py | 2 +- src/scvi/module/_vae.py | 179 +------------------------ src/scvi/train/_trainingplans.py | 12 +- 6 files changed, 26 insertions(+), 190 deletions(-) diff --git a/src/scvi/external/assayvi/_model.py b/src/scvi/external/assayvi/_model.py index 2e8d0fa261..8630a2c959 100644 --- a/src/scvi/external/assayvi/_model.py +++ b/src/scvi/external/assayvi/_model.py @@ -11,6 +11,7 @@ from scvi.data.fields import ( CategoricalJointObsField, CategoricalObsField, + LabelsWithUnlabeledObsField, LayerField, NumericalJointObsField, ) @@ -130,7 +131,6 @@ def __init__( **kwargs, ): super().__init__(adata) - print('22222211') self._module_kwargs = { "n_hidden": n_hidden, @@ -177,7 +177,7 @@ def __init__( n_input=self.summary_stats.n_vars, n_batch=self.summary_stats.n_batch, n_assay=self.summary_stats.n_assay, - n_labels=self.summary_stats.n_labels, + n_labels=self.summary_stats.get("n_labels", 1), n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), n_cats_per_cov=n_cats_per_cov, n_hidden=n_hidden, @@ -328,6 +328,7 @@ def setup_anndata( batch_key: str | None = None, assay_key: str | None = None, labels_key: str | None = None, + unlabeled_category: str = "unlabeled", categorical_covariate_keys: list[str] | None = None, continuous_covariate_keys: list[str] | None = None, **kwargs, @@ -341,7 +342,8 @@ def setup_anndata( %(param_batch_key)s assay_key Key in ``adata.obs`` that corresponds to the assay of the data. - %(param_label_key)s + %(param_labels_key)s + %(param_unlabeled_category)s %(param_cat_cov_keys)s %(param_cont_cov_keys)s """ @@ -350,10 +352,13 @@ def setup_anndata( LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), - CategoricalObsField(REGISTRY_KEYS.LABELS_KEY, labels_key), CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), NumericalJointObsField(REGISTRY_KEYS.CONT_COVS_KEY, continuous_covariate_keys), ] + if labels_key is not None: + anndata_fields.append( + LabelsWithUnlabeledObsField( + REGISTRY_KEYS.LABELS_KEY, labels_key, unlabeled_category)) # register new fields if the adata is minified adata_minify_type = _get_adata_minify_type(adata) if adata_minify_type is not None: diff --git a/src/scvi/external/assayvi/_module.py b/src/scvi/external/assayvi/_module.py index 391c2026e5..59742d44d5 100644 --- a/src/scvi/external/assayvi/_module.py +++ b/src/scvi/external/assayvi/_module.py @@ -528,7 +528,7 @@ def loss( from torch.distributions import kl_divergence x = tensors[REGISTRY_KEYS.X_KEY] - y = tensors[REGISTRY_KEYS.LABELS_KEY] + y = tensors.get(REGISTRY_KEYS.LABELS_KEY, None) kl_divergence_z = self.prior.kl( qz=inference_outputs[MODULE_KEYS.QZ_KEY], diff --git a/src/scvi/model/_scvi.py b/src/scvi/model/_scvi.py index fb9ee4e59b..d527dfd554 100644 --- a/src/scvi/model/_scvi.py +++ b/src/scvi/model/_scvi.py @@ -180,10 +180,7 @@ def __init__( n_cats_per_cov = None n_batch = self.summary_stats.n_batch - n_assay = self.summary_stats.n_assay - use_size_factor_key = self.registry_["setup_args"][ - f"{REGISTRY_KEYS.SIZE_FACTOR_KEY}_key" - ] + use_size_factor_key = REGISTRY_KEYS.SIZE_FACTOR_KEY in self.adata_manager.data_registry library_log_means, library_log_vars = None, None if ( not use_size_factor_key @@ -196,7 +193,6 @@ def __init__( self.module = self._module_cls( n_input=self.summary_stats.n_vars, n_batch=n_batch, - n_assay=n_assay, n_labels=self.summary_stats.n_labels, n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), n_cats_per_cov=n_cats_per_cov, @@ -224,7 +220,6 @@ def setup_anndata( adata: AnnData, layer: str | None = None, batch_key: str | None = None, - assay_key: str | None = None, labels_key: str | None = None, size_factor_key: str | None = None, categorical_covariate_keys: list[str] | None = None, @@ -247,7 +242,6 @@ def setup_anndata( anndata_fields = [ LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), - CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), CategoricalObsField(REGISTRY_KEYS.LABELS_KEY, labels_key), NumericalObsField(REGISTRY_KEYS.SIZE_FACTOR_KEY, size_factor_key, required=False), CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), diff --git a/src/scvi/model/base/_training_mixin.py b/src/scvi/model/base/_training_mixin.py index 85b2671669..52d09cc054 100644 --- a/src/scvi/model/base/_training_mixin.py +++ b/src/scvi/model/base/_training_mixin.py @@ -41,7 +41,7 @@ class UnsupervisedTrainingMixin: """General purpose unsupervised train method.""" _data_splitter_cls = DataSplitter - _training_plan_cls = AdversarialTrainingPlan + _training_plan_cls = TrainingPlan _train_runner_cls = TrainRunner @devices_dsp.dedent diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index 1eda3831ef..0e7f9c88c6 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -173,10 +173,6 @@ def __init__( extra_encoder_kwargs: dict | None = None, extra_decoder_kwargs: dict | None = None, batch_embedding_kwargs: dict | None = None, - conditional_norm: dict | None = None, - mmd_kernel: str | None = "rbf", - prior: str | None = None, - num_classes: int | None = 30, ): from scvi.nn import DecoderSCVI, Encoder @@ -194,8 +190,6 @@ def __init__( 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 - self.mmd_kernel = mmd_kernel if not self.use_observed_lib_size: if library_log_means is None or library_log_vars is None: @@ -254,7 +248,6 @@ def __init__( use_layer_norm=use_layer_norm_encoder, var_activation=var_activation, return_dist=True, - conditional_norm=conditional_norm, **_extra_encoder_kwargs, ) # l encoder goes from n_input-dimensional data to 1-d library size @@ -290,20 +283,6 @@ def __init__( **_extra_decoder_kwargs, ) - self.prior = prior - if prior == "mog": - self.register_parameter( - "prior_means", - torch.nn.Parameter(torch.randn([num_classes, n_latent])), - ) - self.register_parameter( - "prior_log_scales", - torch.nn.Parameter(torch.zeros([num_classes, n_latent])), - ) - self.register_parameter( - "prior_logits", torch.nn.Parameter(torch.ones([num_classes])) - ) - def _get_inference_input( self, tensors: dict[str, torch.Tensor | None], @@ -325,7 +304,6 @@ def _get_inference_input( MODULE_KEYS.X_KEY: tensors[REGISTRY_KEYS.X_KEY], MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY], MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY], - MODULE_KEYS.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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), } @@ -350,7 +328,6 @@ def _get_generative_input( 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.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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), @@ -385,7 +362,6 @@ def _regular_inference( self, x: torch.Tensor, batch_index: torch.Tensor, - assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, n_samples: int = 1, @@ -407,10 +383,7 @@ def _regular_inference( else: categorical_input = () - if assay_index is not None: - assay_index = assay_index.long() - qz, z = self.z_encoder(encoder_input, assay_index, *categorical_input) - elif self.encode_covariates and self.batch_representation == "embedding": + if self.encode_covariates and self.batch_representation == "embedding": 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) @@ -476,7 +449,6 @@ def generative( z: torch.Tensor, library: torch.Tensor, batch_index: torch.Tensor, - assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, size_factor: torch.Tensor | None = None, @@ -524,7 +496,6 @@ def generative( size_factor, *categorical_input, y, - assay=assay_index.long() if assay_index is not None else None, ) else: px_scale, px_r, px_rate, px_dropout = self.decoder( @@ -534,7 +505,6 @@ def generative( batch_index, *categorical_input, y, - assay=assay_index.long() if assay_index is not None else None, ) if self.dispersion == "gene-label": @@ -585,28 +555,15 @@ def loss( 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, - weight_assay_loss: float = 1.0, + kl_weight: float = 1.0, ) -> LossOutput: """Compute the loss.""" from torch.distributions import kl_divergence x = tensors[REGISTRY_KEYS.X_KEY] - if self.prior == "mog": - qz = inference_outputs[MODULE_KEYS.QZ_KEY] - cats = distributions.Categorical(logits=self.prior_logits) - normal_dists = distributions.Independent( - distributions.Normal(self.prior_means, torch.exp(self.prior_log_scales) + 1e-4), - 1, - ) - prior = distributions.MixtureSameFamily(cats, normal_dists) - u = qz.rsample(sample_shape=(30,)) - # (sample, n_obs, n_latent) -> (sample, n_obs,) - kl_divergence_z = (qz.log_prob(u).sum(-1) - prior.log_prob(u)).mean(0) - else: - kl_divergence_z = kl_divergence( - inference_outputs[MODULE_KEYS.QZ_KEY], generative_outputs[MODULE_KEYS.PZ_KEY] - ).sum(dim=-1) + 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] @@ -614,15 +571,6 @@ def loss( else: kl_divergence_l = torch.zeros_like(kl_divergence_z) - assay_index = tensors.get(REGISTRY_KEYS.ASSAY_KEY, None) - if weight_assay_loss > 0.0 and assay_index is not None: - assay_loss = self._compute_assay_penalty( - inference_outputs[MODULE_KEYS.QZ_KEY].loc, - assay_index - ) - else: - assay_loss = 0.0 - reconst_loss = -generative_outputs[MODULE_KEYS.PX_KEY].log_prob(x).sum(-1) kl_local_for_warmup = kl_divergence_z @@ -630,7 +578,7 @@ def loss( weighted_kl_local = kl_weight * kl_local_for_warmup + kl_local_no_warmup - loss = torch.mean(reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss) + loss = torch.mean(reconst_loss + weighted_kl_local) # a payload to be used during autotune if self.extra_payload_autotune: @@ -799,121 +747,6 @@ def marginal_ll( batch_log_lkl = batch_log_lkl.cpu() return batch_log_lkl - def _compute_assay_penalty( - self, params, assay): - assay = assay.ravel().long() - unique = torch.unique(assay) - pair_penalty = torch.tensor(0., device=assay.device) - if len(unique) > 1: - for i in unique: - pp = self.mmd(params, mask=(assay == i)) - pair_penalty += pp - - return pair_penalty - - def mmd(self, params, mask=None): - if mask is not None: - mod_1 = params[mask] - mod_2 = params[~mask] - if self.mmd_kernel == 'imq': - penalty = imq_kernel(mod_1, mod_2, beta=0.5) - elif self.mmd_kernel == 'rbf': - penalty = rbf_kernel(mod_1, mod_2) - else: - penalty = torch.linalg.norm(mod_1 - mod_2, dim=1).mean() - return penalty - -def imq_kernel(x, y, gammas=None, beta=0.5): - """ - Compute the IMQ kernel between two tensors. - - Parameters - ---------- - x : torch.Tensor - Input tensor of shape (N, D). - y : torch.Tensor - Input tensor of shape (M, D). - gammas : list of float or None - List of gamma values to compute the kernel with. - beta : float - The beta parameter controlling the sharpness of the kernel. - - Returns - ------- - kernel_sum : torch.Tensor - Tensor of shape (N, M), the sum of IMQ kernels over all gamma values. - """ - if gammas is None: - gammas = [ - 1e-3, - 1e-2, - 1e-1, - 1, - 5, - 10, - ] - kxy = torch.cdist(x, y).pow(2) - kxx = torch.cdist(x, x).pow(2) - kyy = torch.cdist(y, y).pow(2) - - kernel_sum_xy = torch.tensor(0.0, device=x.device) - kernel_sum_xx = torch.tensor(0.0, device=x.device) - kernel_sum_yy = torch.tensor(0.0, device=x.device) - for gamma in gammas: - kernel_sum_xy += (gamma + kxy).pow(-beta).mean() - kernel_sum_xx += (gamma + kxx).pow(-beta).mean() - kernel_sum_yy += (gamma + kyy).pow(-beta).mean() - - return (kernel_sum_xx + kernel_sum_yy - 2 * kernel_sum_xy) / len(gammas) - -@auto_move_data -def rbf_kernel(x, y, gammas=None): - """ - Compute the RBF kernel between two tensors. - - Parameters - ---------- - x : torch.Tensor - Input tensor of shape (N, D). - y : torch.Tensor - Input tensor of shape (M, D). - gammas : list of float or None - List of gamma values to compute the kernel with. - - Returns - ------- - kernel_sum : torch.Tensor - Tensor of shape (N, M), the sum of RBF kernels over all gamma values. - """ - if gammas is None: - gammas = [ - 1e-10, - 1e-8, - 1e-6, - 1e-4, - 1e-3, - 1e-2, - 1e-1, - 1, - 2, - 5, - 10, - ] - kxy = torch.cdist(x, y).pow(2) - kxx = torch.cdist(x, x).pow(2) - kyy = torch.cdist(y, y).pow(2) - - kernel_sum_xy = torch.tensor(0.0, device=x.device) - kernel_sum_xx = torch.tensor(0.0, device=x.device) - kernel_sum_yy = torch.tensor(0.0, device=x.device) - for gamma in gammas: - kernel_sum_xy += torch.exp(-gamma * kxy).mean() - kernel_sum_xx += torch.exp(-gamma * kxx).mean() - kernel_sum_yy += torch.exp(-gamma * kyy).mean() - #print(gamma, (torch.exp(-gamma * kxx).mean() + torch.exp(-gamma * kyy).mean() - 2 * torch.exp(-gamma * kxy).mean())) - - return (kernel_sum_xx + kernel_sum_yy - 2 * kernel_sum_xy) / len(gammas) - class LDVAE(VAE): """Linear-decoded Variational auto-encoder model. diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 59ee13517a..762260c533 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -725,17 +725,21 @@ def training_step(self, batch, batch_idx): # train adversarial classifier # this condition will not be met unless self.adversarial_classifier is not False if opt2 is not None: - for _ in range(self.adversarial_steps): + loss = 0. + for i in range(self.adversarial_steps): qz = inference_outputs["qz"] z = qz.sample().detach() - loss = self.loss_adversarial_classifier(z, batch_tensor, True) + loss_ = kappa * self.loss_adversarial_classifier(z, batch_tensor, True) if isinstance(self.module.prior, MogPrior) or isinstance(self.module.prior, VampPrior): qz_m, qz_v = qz.loc.detach(), qz.scale.detach() - loss += self.module.prior.kl( + loss_ += self.module.prior.kl( qz=Normal(qz_m, qz_v), z=z, - labels=batch[REGISTRY_KEYS.LABELS_KEY].long(), + labels=batch.get(REGISTRY_KEYS.LABELS_KEY, torch.tensor(0)).long(), ).mean() + if i>1 and (loss - loss_)/loss < 1e-3: + break + loss = loss_ opt2.zero_grad() self.manual_backward(loss) opt2.step() From ff92f2f257888545815e8ccca86cae81b2af9f6d Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 28 May 2025 16:18:17 -0700 Subject: [PATCH 05/24] Undo changes --- src/scvi/external/assayvi/__init__.py | 1 + src/scvi/model/base/_archesmixin.py | 1 - 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/src/scvi/external/assayvi/__init__.py b/src/scvi/external/assayvi/__init__.py index e7719ae44e..c6ca875016 100644 --- a/src/scvi/external/assayvi/__init__.py +++ b/src/scvi/external/assayvi/__init__.py @@ -2,3 +2,4 @@ from ._module import ASSAYVAE __all__ = ["ASSAYVI", "ASSAYVAE"] +#323d5 \ No newline at end of file diff --git a/src/scvi/model/base/_archesmixin.py b/src/scvi/model/base/_archesmixin.py index 64222ea206..a098b63847 100644 --- a/src/scvi/model/base/_archesmixin.py +++ b/src/scvi/model/base/_archesmixin.py @@ -414,7 +414,6 @@ def _set_params_online_update( if not freeze_classifier: mod_no_hooks_yes_grad.add("classifier") parameters_yes_grad = {"background_pro_alpha", "background_pro_log_beta"} - parameters_yes_grad = {"background_pro_alpha", "background_pro_log_beta"} def no_hook_cond(key): one = (not freeze_expression) and "encoder" in key From 05f9b92dd27561d42910365fd52d3776404f1f8a Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 28 May 2025 16:19:36 -0700 Subject: [PATCH 06/24] bdb --- src/scvi/external/assayvi/_model.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/scvi/external/assayvi/_model.py b/src/scvi/external/assayvi/_model.py index 8630a2c959..9b2c8f9063 100644 --- a/src/scvi/external/assayvi/_model.py +++ b/src/scvi/external/assayvi/_model.py @@ -39,6 +39,7 @@ from anndata import AnnData logger = logging.getLogger(__name__) +print(2) class ASSAYVI( From d165a0b9de185bcb7b34ac58029ba38e141772b9 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 28 May 2025 16:42:01 -0700 Subject: [PATCH 07/24] Fix name GaussianPrior --- src/scvi/external/assayvi/__init__.py | 3 +-- src/scvi/external/assayvi/_module.py | 4 ++-- src/scvi/external/sysvi/_module.py | 4 ++-- src/scvi/external/sysvi/_priors.py | 2 +- 4 files changed, 6 insertions(+), 7 deletions(-) diff --git a/src/scvi/external/assayvi/__init__.py b/src/scvi/external/assayvi/__init__.py index c6ca875016..9ccdbfa0f0 100644 --- a/src/scvi/external/assayvi/__init__.py +++ b/src/scvi/external/assayvi/__init__.py @@ -1,5 +1,4 @@ from ._model import ASSAYVI from ._module import ASSAYVAE -__all__ = ["ASSAYVI", "ASSAYVAE"] -#323d5 \ No newline at end of file +__all__ = ["ASSAYVI", "ASSAYVAE"] \ No newline at end of file diff --git a/src/scvi/external/assayvi/_module.py b/src/scvi/external/assayvi/_module.py index 59742d44d5..d4aaa1acbb 100644 --- a/src/scvi/external/assayvi/_module.py +++ b/src/scvi/external/assayvi/_module.py @@ -15,8 +15,8 @@ BaseMinifiedModeModuleClass, EmbeddingModuleMixin, LossOutput, + GaussianPrior, MogPrior, - StandardPrior, VampPrior, auto_move_data, ) @@ -285,7 +285,7 @@ def __init__( **cls_parameters, ) if prior == "gaussian": - self.prior = StandardPrior() + self.prior = GaussianPrior() elif prior == "vamp": assert pseudoinput_data is not None, ( "Pseudoinput data must be specified if using VampPrior" diff --git a/src/scvi/external/sysvi/_module.py b/src/scvi/external/sysvi/_module.py index b06857e439..353ace8f2b 100644 --- a/src/scvi/external/sysvi/_module.py +++ b/src/scvi/external/sysvi/_module.py @@ -9,7 +9,7 @@ from scvi.module.base import BaseModuleClass, EmbeddingModuleMixin, LossOutput, auto_move_data from ._base_components import EncoderDecoder -from ._priors import StandardPrior, VampPrior +from ._priors import GaussianPrior, VampPrior if TYPE_CHECKING: from typing import Literal @@ -137,7 +137,7 @@ def __init__( ) if prior == "standard_normal": - self.prior = StandardPrior() + self.prior = GaussianPrior() elif prior == "vamp": assert pseudoinput_data is not None, ( "Pseudoinput data must be specified if using VampPrior" diff --git a/src/scvi/external/sysvi/_priors.py b/src/scvi/external/sysvi/_priors.py index c966fafb9e..89d615f64e 100644 --- a/src/scvi/external/sysvi/_priors.py +++ b/src/scvi/external/sysvi/_priors.py @@ -35,7 +35,7 @@ def kl( pass -class StandardPrior(Prior): +class GaussianPrior(Prior): """Standard prior distribution.""" def kl(self, qz: torch.Tensor, z: None = None) -> torch.Tensor: From ff52ce742cd44b7e2f5e61d0c12bf70f7d474df6 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 18 Jun 2025 01:54:53 -0700 Subject: [PATCH 08/24] Final scvi-x version --- src/scvi/external/__init__.py | 5 ++-- src/scvi/external/assayvi/__init__.py | 4 --- src/scvi/external/scvix/__init__.py | 4 +++ .../external/{assayvi => scvix}/_model.py | 10 +++---- .../external/{assayvi => scvix}/_module.py | 30 +++++++++---------- src/scvi/external/sysvi/_base_components.py | 4 +-- src/scvi/external/sysvi/_module.py | 8 ++--- src/scvi/module/_vae.py | 1 - src/scvi/nn/_base_components.py | 17 +++++------ src/scvi/train/_trainingplans.py | 2 ++ 10 files changed, 42 insertions(+), 43 deletions(-) delete mode 100644 src/scvi/external/assayvi/__init__.py create mode 100644 src/scvi/external/scvix/__init__.py rename src/scvi/external/{assayvi => scvix}/_model.py (98%) rename src/scvi/external/{assayvi => scvix}/_module.py (97%) diff --git a/src/scvi/external/__init__.py b/src/scvi/external/__init__.py index 558af1ccdb..bf642bcce0 100644 --- a/src/scvi/external/__init__.py +++ b/src/scvi/external/__init__.py @@ -2,8 +2,6 @@ from scvi import settings from scvi.utils import error_on_missing_dependencies - -from .assayvi import ASSAYVI from .cellassign import CellAssign from .contrastivevi import ContrastiveVI from .cytovi import CYTOVI @@ -19,6 +17,7 @@ from .scar import SCAR from .scbasset import SCBASSET from .scviva import SCVIVA +from .scvix import SCVIX from .solo import SOLO from .stereoscope import RNAStereoscope, SpatialStereoscope from .sysvi import SysVI @@ -26,7 +25,7 @@ from .velovi import VELOVI __all__ = [ - "ASSAYVI", + "SCVIX", "SCAR", "SOLO", "GIMVI", diff --git a/src/scvi/external/assayvi/__init__.py b/src/scvi/external/assayvi/__init__.py deleted file mode 100644 index 9ccdbfa0f0..0000000000 --- a/src/scvi/external/assayvi/__init__.py +++ /dev/null @@ -1,4 +0,0 @@ -from ._model import ASSAYVI -from ._module import ASSAYVAE - -__all__ = ["ASSAYVI", "ASSAYVAE"] \ No newline at end of file diff --git a/src/scvi/external/scvix/__init__.py b/src/scvi/external/scvix/__init__.py new file mode 100644 index 0000000000..aea7c60506 --- /dev/null +++ b/src/scvi/external/scvix/__init__.py @@ -0,0 +1,4 @@ +from ._model import SCVIX +from ._module import VAEX + +__all__ = ["SCVIX", "VAEX"] \ No newline at end of file diff --git a/src/scvi/external/assayvi/_model.py b/src/scvi/external/scvix/_model.py similarity index 98% rename from src/scvi/external/assayvi/_model.py rename to src/scvi/external/scvix/_model.py index 9b2c8f9063..b5ab29ae4e 100644 --- a/src/scvi/external/assayvi/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -30,7 +30,7 @@ from scvi.utils import setup_anndata_dsp from scvi.utils._docstrings import devices_dsp, setup_anndata_dsp -from ._module import ASSAYVAE +from ._module import VAEX if TYPE_CHECKING: from typing import Literal @@ -42,7 +42,7 @@ print(2) -class ASSAYVI( +class SCVIX( EmbeddingMixin, RNASeqMixin, VAEMixin, @@ -110,9 +110,9 @@ class ASSAYVI( :class:`~scvi.module.VAE` """ - _module_cls = ASSAYVAE - _LATENT_QZM_KEY = "assayvi_latent_qzm" - _LATENT_QZV_KEY = "assayvi_latent_qzv" + _module_cls = VAEX + _LATENT_QZM_KEY = "scvix_latent_qzm" + _LATENT_QZV_KEY = "scvix_latent_qzv" _data_splitter_cls = DataSplitter _training_plan_cls = AdversarialTrainingPlan _train_runner_cls = TrainRunner diff --git a/src/scvi/external/assayvi/_module.py b/src/scvi/external/scvix/_module.py similarity index 97% rename from src/scvi/external/assayvi/_module.py rename to src/scvi/external/scvix/_module.py index d4aaa1acbb..b86ab11fa7 100644 --- a/src/scvi/external/assayvi/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -31,7 +31,7 @@ logger = logging.getLogger(__name__) -class ASSAYVAE(EmbeddingModuleMixin, BaseMinifiedModeModuleClass): +class VAEX(EmbeddingModuleMixin, BaseMinifiedModeModuleClass): """Variational auto-encoder :cite:p:`Lopez18`. Parameters @@ -162,7 +162,6 @@ def __init__( prior: str | None = None, pseudoinput_data: dict | None = None, n_prior_components: int | None = 30, - mmd_kernel: str = "rbf", ): from scvi.nn import DecoderSCVI, Encoder @@ -178,6 +177,7 @@ def __init__( self.latent_distribution = latent_distribution self.encode_covariates = encode_covariates self.use_observed_lib_size = True + self.n_hidden = n_hidden if self.dispersion == "gene": @@ -201,7 +201,7 @@ def __init__( self.init_embedding(REGISTRY_KEYS.BATCH_KEY, n_batch, **(batch_embedding_kwargs or {})) n_continuous += self.get_embedding_dim(REGISTRY_KEYS.BATCH_KEY) elif self.batch_representation != "one-hot": - raise ValueError("`batch_representation` must be one of 'one-hot', 'embedding'.") + raise ValueError("`batch_representation` must be one of 'one-hot' or '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" @@ -213,7 +213,6 @@ def __init__( else: cat_list = [n_batch] + n_cats_per_cov_ self.encode_assay = encode_assay - self.batch_representation_encoder = False conditional_category = 0 if encode_assay: encode_assay_list = [n_assay] @@ -229,6 +228,7 @@ def __init__( encoder_cat_list = encode_assay_list + n_cats_per_cov_ n_cont_encoder = n_continuous else: + self.batch_representation_encoder = False encoder_cat_list = encode_assay_list + [n_batch] + n_cats_per_cov_ n_cont_encoder = n_continuous_cov if conditional_norm and not encode_assay: @@ -294,7 +294,6 @@ def __init__( pseudoinput_data, full_forward_pass=True ) - print('include training.') cat_list = [n_batch] + n_cats_per_cov_ + encode_assay_list self.prior = VampPrior( n_components=n_prior_components, @@ -391,17 +390,18 @@ def _regular_inference( if self.encode_covariates and self.batch_representation_encoder: batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index) - if cont_covs is not None: - cont_input = torch.cat([cont_covs, batch_rep], dim=-1) - else: - cont_input = batch_rep + if cont_covs is not None and self.encode_covariates: + cont = torch.cat([cont_covs, batch_rep], dim=-1) else: - cont_input = cont_covs + if self.encode_covariates: + cont = batch_rep + else: + cont = None if not self.encode_assay: assay_index = None else: assay_index = assay_index.long() - qz, z = self.z_encoder(x_, assay_index, batch_index, *categorical_input, cont_input=cont_input) + qz, z = self.z_encoder(x_, assay_index, batch_index, *categorical_input, cont=cont) if n_samples > 1: untran_z = qz.sample((n_samples,)) @@ -453,11 +453,11 @@ def generative( if self.batch_representation == "embedding": batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index) if cont_covs is not None: - cont_input = torch.cat([cont_covs, batch_rep], dim=-1) + cont = torch.cat([cont_covs, batch_rep], dim=-1) else: - cont_input = batch_rep + cont = batch_rep else: - cont_input = cont_covs + cont = cont_covs px_scale, px_r, px_rate, px_dropout = self.decoder( self.dispersion, @@ -466,7 +466,7 @@ def generative( batch_index, *categorical_input, y, - cont_input=cont_input, + cont=cont, output_condition=assay_index.long(), ) diff --git a/src/scvi/external/sysvi/_base_components.py b/src/scvi/external/sysvi/_base_components.py index 81f8d1a859..98430d3beb 100644 --- a/src/scvi/external/sysvi/_base_components.py +++ b/src/scvi/external/sysvi/_base_components.py @@ -59,7 +59,7 @@ def __init__( n_input: int, n_output: int, n_cat_list: list[int], - n_cont: int, + n_continuous: int, n_hidden: int = 256, n_layers: int = 3, var_mode: Literal["sample_feature", "feature"] = "feature", @@ -73,7 +73,7 @@ def __init__( self.decoder_y = FCLayers( n_in=n_input, n_cat_list=n_cat_list, - n_cont=n_cont, + n_continuous=n_continuous, n_out=n_hidden, n_hidden=n_hidden, n_layers=n_layers, diff --git a/src/scvi/external/sysvi/_module.py b/src/scvi/external/sysvi/_module.py index 353ace8f2b..9712845476 100644 --- a/src/scvi/external/sysvi/_module.py +++ b/src/scvi/external/sysvi/_module.py @@ -100,13 +100,13 @@ def __init__( self.n_batch = n_batch n_cat_list = [n_batch] - n_cont = n_continuous_cov + n_continuous = n_continuous_cov if n_cats_per_cov is not None: if self.embed_categorical_covariates: for idx, n in enumerate(n_cats_per_cov): covariate_name = f"cov{idx}" self.init_embedding(covariate_name, n, **embedding_kwargs) - n_cont += self.get_embedding(covariate_name).embedding_dim + n_continuous += self.get_embedding(covariate_name).embedding_dim else: n_cat_list.extend(n_cats_per_cov) @@ -114,7 +114,7 @@ def __init__( n_input=n_input, n_output=n_latent, n_cat_list=n_cat_list, - n_cont=n_cont, + n_continuous=n_continuous, n_hidden=n_hidden, n_layers=n_layers, dropout_rate=dropout_rate, @@ -127,7 +127,7 @@ def __init__( n_input=n_latent, n_output=n_input, n_cat_list=n_cat_list, - n_cont=n_cont, + n_continuous=n_continuous, n_hidden=n_hidden, n_layers=n_layers, dropout_rate=dropout_rate, diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index 0e7f9c88c6..fae77d8242 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -7,7 +7,6 @@ import numpy as np import torch from torch.nn.functional import one_hot -from torch import distributions from scvi import REGISTRY_KEYS, settings from scvi.data._constants import ADATA_MINIFY_TYPE diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index 6611790723..cf1ede94d1 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -90,7 +90,6 @@ def __init__( n_out: int, n_continuous: int = 0, n_cat_list: Iterable[int] = None, - n_cont: int = 0, n_layers: int = 1, n_hidden: int = 128, dropout_rate: float = 0.1, @@ -116,7 +115,7 @@ def __init__( self.cond_cat = conditional_category - self.n_cov = n_cont + sum(self.n_cat_list) + self.n_cov = n_continuous + sum(self.n_cat_list) self.fc_layers = nn.Sequential( collections.OrderedDict( @@ -226,7 +225,7 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No layer, ConditionalLayerNorm): if x.dim() == 3: x = torch.cat( - [(layer(slice_x, cat_list[self.cond_cat])).unsqueeze(0) + [(layer(x=slice_x, y=cat_list[self.cond_cat])).unsqueeze(0) for slice_x in x], dim=0 ) @@ -342,7 +341,7 @@ def forward( self, x: torch.Tensor, *cat_list: int, - cont_input: torch.Tensor | None = None, + cont: torch.Tensor | None = None, ): r"""The forward computation for a single sample. @@ -357,7 +356,7 @@ def forward( tensor with shape (n_input,) cat_list list of category membership(s) for this sample - cont_input + cont optional tensor with shape (n_continuous,) Returns @@ -367,7 +366,7 @@ def forward( """ # Parameters for latent distribution - q = self.encoder(x, *cat_list, cont_input=cont_input) + q = self.encoder(x, *cat_list, cont=cont) q_m = self.mean_encoder(q) q_v = self.var_activation(self.var_encoder(q)) + self.var_eps dist = Normal(q_m, q_v.sqrt()) @@ -470,7 +469,7 @@ def forward( z: torch.Tensor, library: torch.Tensor, *cat_list: int, - cont_input: torch.Tensor | None = None, + cont: torch.Tensor | None = None, output_condition: torch.Tensor | None = None, ): """The forward computation for a single sample. @@ -494,7 +493,7 @@ def forward( library size cat_list list of category membership(s) for this sample - cont_input + cont tensor with shape ``(n_continuous,)`` output_condition tensor with shape ``(n_input,)`` used for conditioning the output layer @@ -506,7 +505,7 @@ def forward( """ # The decoder returns values for the parameters of the ZINB distribution - px = self.px_decoder(z, *cat_list, cont_input=cont_input) + px = self.px_decoder(z, *cat_list, cont=cont) if output_condition is not None and self.n_conditions_output: one_hot_cat = nn.functional.one_hot( output_condition.squeeze(-1), diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 762260c533..cff62bb66c 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -579,6 +579,8 @@ class AdversarialTrainingPlan(TrainingPlan): Whether to use adversarial classifier in the latent space adversarial_key Key in setup args to use for adversarial training. + adversarial_steps + Number of steps to train the adversarial classifier for each training step. scale_adversarial_loss Scaling factor on the adversarial components of the loss. By default, adversarial loss is scaled from 1 to 0 following the opposite of From 7cb54ecd448ddc17d0bed7e1e1539b304abbd9fa Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Tue, 8 Jul 2025 03:46:56 -0700 Subject: [PATCH 09/24] adding da test --- src/scvi/model/base/_da_testing.py | 116 +++++++++++++++++++++++++++++ 1 file changed, 116 insertions(+) create mode 100644 src/scvi/model/base/_da_testing.py diff --git a/src/scvi/model/base/_da_testing.py b/src/scvi/model/base/_da_testing.py new file mode 100644 index 0000000000..e50e701153 --- /dev/null +++ b/src/scvi/model/base/_da_testing.py @@ -0,0 +1,116 @@ +from collections.abc import Sequence + +import numpy as np +import pandas as pd +import torch +import torch.distributions as dist +from anndata import AnnData +from tqdm import tqdm + + +def get_aggregated_posterior( + self, + adata: AnnData | None = None, + sample: str | int | None = None, + indices: Sequence[int] | None = None, + batch_size: int | None = None, + dof: float | None = 3., +) -> dist.Distribution: + """Compute the aggregated posterior over the ``u`` latent representations. + + Parameters + ---------- + adata + AnnData object to use. Defaults to the AnnData object used to initialize the model. + sample + Name or index of the sample to filter on. If ``None``, uses all cells. + indices + Indices of cells to use. + batch_size + Batch size to use for computing the latent representation. + dof + Degrees of freedom for the Student's t-distribution components. If ``None``, components are Normal. + + Returns + ------- + A mixture distribution of the aggregated posterior. + """ + self._check_if_trained(warn=False) + adata = self._validate_anndata(adata) + + if indices is None: + indices = np.arange(self.adata.n_obs) + if sample is not None: + indices = np.intersect1d( + np.array(indices), np.where(adata.obs[self.sample_key] == sample)[0] + ) + + dataloader = self._make_data_loader(adata=adata, indices=indices, batch_size=batch_size) + qu_loc, qu_scale = self.get_latent_representation(batch_size=batch_size, return_dist=True, dataloader=dataloader, give_mean=True) + + qu_loc = torch.tensor(qu_loc, device='cuda').T + qu_scale = torch.tensor(qu_scale, device='cuda').T + + if dof is None: + components = dist.Normal(qu_loc, qu_scale) + else: + components = dist.StudentT(dof, qu_loc, qu_scale) + return dist.MixtureSameFamily( + dist.Categorical(logits=torch.ones(qu_loc.shape[1], device='cuda')), components) + +def differential_abundance( + self, + adata: AnnData | None = None, + sample_key: str | None = None, + batch_size: int = 128, + downsample_cells: int | None = None, + dof: float | None = None, +) -> pd.DataFrame: + """Compute the differential abundance between samples. + + Computes the logarithm of the ratio of the probabilities of each sample conditioned on the + estimated aggregate posterior distribution of each cell. + + Parameters + ---------- + adata + The data object to compute the differential abundance for. + sample_key + Key for the sample covariate. + batch_size + Minibatch size for computing the differential abundance. + downsample_cells + Number of cells to subset to before computing the differential abundance. + dof + Degrees of freedom for the Student's t-distribution components for aggregated posterior. If ``None``, components are Normal. + + Returns + ------- + DataFrame of shape (n_cells, n_samples) containing the log probabilities + for each cell across samples. The rows correspond to cell names from `adata.obs_names`, + and the columns correspond to unique sample identifiers. + """ + adata = self._validate_anndata(adata) + + us = self.get_latent_representation( + batch_size=batch_size, return_dist=False, give_mean=True + ) + + unique_samples = adata.obs[sample_key].unique() + dataloader = torch.utils.data.DataLoader(us, batch_size=batch_size) + log_probs = [] + for sample_name in tqdm(unique_samples): + indices = np.where(adata.obs[sample_key] == sample_name)[0] + if downsample_cells is not None and downsample_cells < indices.shape[0]: + indices = np.random.choice(indices, downsample_cells, replace=False) + + ap = get_aggregated_posterior(self, adata=adata, indices=indices, dof=dof) + log_probs_ = [] + for u_rep in dataloader: + u_rep = u_rep.to('cuda') + log_probs_.append(ap.log_prob(u_rep).sum(-1, keepdims=True)) + log_probs.append(torch.cat(log_probs_, axis=0).cpu().numpy()) + + log_probs = np.concatenate(log_probs, 1) + log_probs_df = pd.DataFrame(data=log_probs, index=adata.obs_names.to_numpy(), columns=unique_samples) + return log_probs_df From 3923491ee9ad1c511c592fe478acd65ad4c39cc1 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Thu, 10 Jul 2025 03:05:52 -0700 Subject: [PATCH 10/24] Updates scArches and minified --- src/scvi/external/scvix/_model.py | 18 +++- src/scvi/external/scvix/_module.py | 37 +++++-- src/scvi/model/base/_archesmixin.py | 3 +- src/scvi/nn/_base_components.py | 4 +- tests/external/scvix/test_scvix.py | 152 ++++++++++++++++++++++++++++ 5 files changed, 196 insertions(+), 18 deletions(-) create mode 100644 tests/external/scvix/test_scvix.py diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index b5ab29ae4e..efca6dfc35 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -79,11 +79,19 @@ class SCVIX( * ``'zinb'`` - Zero-inflated negative binomial distribution * ``'poisson'`` - Poisson distribution * ``'normal'`` - ``EXPERIMENTAL`` Normal distribution - latent_distribution + prior One of: - - * ``'normal'`` - Normal distribution - * ``'ln'`` - Logistic normal distribution (Normal(0, I) transformed by softmax) + * ``'gaussian'`` - Gaussian prior + * ``'mog'`` - Mixture of Gaussians prior + * ``'vamp'`` - Variational Amortized Mixture of Posteriors prior + * ``'mog_celltype'`` - Mixture of Gaussians prior with cell-type bias + pseudoinputs_data_indices + Indices of cells to use as pseudoinputs for the VAMP prior. If ``None``, a random sample of + ``n_prior_components`` cells will be used. + n_prior_components + Number of components to use for the VAMP and MOG priors. This is the number of pseudoinputs + used for the VAMP prior, and the number of components in the MOG prior. + Defaults to 50. **kwargs Additional keyword arguments for :class:`~scvi.module.VAE`. @@ -126,7 +134,7 @@ def __init__( dropout_rate: float = 0.05, dispersion: Literal["gene", "gene-batch", "gene-cell"] = "gene", gene_likelihood: Literal["zinb", "nb", "poisson", "normal"] = "nb", - prior: Literal["normal", "mog", "vamp"] = "normal", + prior: Literal["gaussian", "mog", "vamp", "mog_celltype"] = "gaussian", pseudoinputs_data_indices: np.array | None = None, n_prior_components: int = 50, **kwargs, diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index b86ab11fa7..58f8c7c5ab 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -74,11 +74,6 @@ class VAEX(EmbeddingModuleMixin, BaseMinifiedModeModuleClass): * ``"zinb"``: :class:`~scvi.distributions.ZeroInflatedNegativeBinomial`. * ``"poisson"``: :class:`~scvi.distributions.Poisson`. * ``"normal"``: :class:`~torch.distributions.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 gene expression is used. @@ -146,7 +141,6 @@ def __init__( dispersion: Literal["gene", "gene-batch", "gene-assay", "gene-cell"] = "gene", log_variational: bool = True, gene_likelihood: Literal["zinb", "nb", "poisson"] = "nb", - latent_distribution: Literal["normal", "ln"] = "normal", encode_covariates: bool = False, encode_assay: bool = True, deeply_inject_covariates: bool = False, @@ -161,7 +155,7 @@ def __init__( conditional_output: bool = True, prior: str | None = None, pseudoinput_data: dict | None = None, - n_prior_components: int | None = 30, + n_prior_components: int | None = 50, ): from scvi.nn import DecoderSCVI, Encoder @@ -174,7 +168,6 @@ def __init__( self.n_batch = n_batch self.n_assay = n_assay self.n_labels = n_labels - self.latent_distribution = latent_distribution self.encode_covariates = encode_covariates self.use_observed_lib_size = True self.n_hidden = n_hidden @@ -242,7 +235,6 @@ def __init__( 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, @@ -394,7 +386,7 @@ def _regular_inference( cont = torch.cat([cont_covs, batch_rep], dim=-1) else: if self.encode_covariates: - cont = batch_rep + cont = cont_covs else: cont = None if not self.encode_assay: @@ -416,6 +408,31 @@ def _regular_inference( 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.LIBRARY_KEY: library, + } + @auto_move_data def generative( self, diff --git a/src/scvi/model/base/_archesmixin.py b/src/scvi/model/base/_archesmixin.py index a098b63847..f9654b4a65 100644 --- a/src/scvi/model/base/_archesmixin.py +++ b/src/scvi/model/base/_archesmixin.py @@ -231,7 +231,6 @@ def load_query_data( freeze_dropout=freeze_dropout, freeze_expression=freeze_expression, freeze_classifier=freeze_classifier, - parameters_yes_grad=additional_parameters, ) model.is_trained_ = False @@ -537,4 +536,4 @@ def _pad_and_sort_query_anndata( if adata_out is not adata: adata._init_as_actual(adata_out) else: - return adata_out + return adata_out \ No newline at end of file diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index cf1ede94d1..dd1a5573e6 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -512,7 +512,9 @@ def forward( self.n_conditions_output ) else: - one_hot_cat = torch.zeros(px.size(0), self.n_conditions_output).to(px.device) + one_hot_cat = torch.zeros(px.size(-2), self.n_conditions_output).to(px.device) + if px.dim() == 3: + one_hot_cat = one_hot_cat.unsqueeze(0).expand(px.size(0), -1, -1) px_cat = torch.cat([px, one_hot_cat], dim=-1) px_scale = self.px_scale_decoder(px_cat) px_dropout = self.px_dropout_decoder(px_cat) diff --git a/tests/external/scvix/test_scvix.py b/tests/external/scvix/test_scvix.py new file mode 100644 index 0000000000..89ec21a764 --- /dev/null +++ b/tests/external/scvix/test_scvix.py @@ -0,0 +1,152 @@ +import pytest +import os +import numpy as np + +import scvi +from scvi.data import synthetic_iid +from scvi.data._constants import ADATA_MINIFY_TYPE +from scvi.data._utils import _is_minified +from scvi.model.base import BaseMinifiedModeModelClass +from scvi.external import SCVIX + + +def assert_approx_equal(a, b): + # Allclose because on GPU, the values are not exactly the same + # as some values are moved to cpu during data minification + np.testing.assert_allclose(a, b, rtol=3e-1, atol=5e-1) + + +@pytest.mark.parametrize("prior", ["gaussian", "mog", "vamp", "mog_celltype"]) +def test_scvix(prior: str): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, prior=prior) + model.train(max_epochs=1) + model.get_latent_representation() + model.get_normalized_expression() + model.get_normalized_expression(transform_batch="batch_1") + model.get_normalized_expression(n_samples=2) + model.get_elbo(indices=model.validation_indices) + model.get_reconstruction_error(indices=model.validation_indices) + model.differential_expression(groupby="labels", group1="label_1") + +@pytest.mark.parametrize("dispersion", ["gene", "gene-batch", "gene-assay", "gene-cell"]) +def test_scvix_dispersion(dispersion: str): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, dispersion=dispersion) + model.train(max_epochs=1) + model.get_normalized_expression() + +def test_scvix_encode_covariates(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, encode_covariates=True) + model.train(max_epochs=1) + model.get_normalized_expression(n_samples=2) + +def test_scvix_embedding(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, batch_representation="embedding") + model.train(max_epochs=1) + model.get_normalized_expression(n_samples=2) + +def test_scvix_layernorm(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, conditional_norm=False, use_batch_norm="both", use_layer_norm="none") + model.train(max_epochs=1) + model.get_normalized_expression(n_samples=2) + model = SCVIX(adata, conditional_norm=True, use_batch_norm="both", use_layer_norm="none") + model.train(max_epochs=1) + model.get_normalized_expression(n_samples=2) + +def test_scvix_scarches_one_hot(save_path): + # test transfer_anndata_setup + view + adata1 = synthetic_iid() + SCVIX.setup_anndata(adata1, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata1, batch_representation="one-hot") + model.train(1, train_size=0.5) + dir_path = os.path.join(save_path, "saved_model/") + model.save(dir_path, overwrite=True) + + # adata2 has more genes and a perfect subset of adata1 + adata2 = synthetic_iid(n_genes=110) + adata2.obs["batch"] = adata2.obs.batch.cat.rename_categories(["batch_2", "batch_3"]) + SCVIX.prepare_query_anndata(adata2, dir_path) + SCVIX_query = SCVIX.load_query_data(adata2, dir_path) + SCVIX_query.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + adata3 = SCVIX.prepare_query_anndata(adata2, dir_path, inplace=False) + SCVIX_query2 = SCVIX.load_query_data(adata3, dir_path) + SCVIX_query2.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + # adata4 has more genes and missing 10 genes from adata1 + adata4 = synthetic_iid(n_genes=110) + new_var_names_init = [f"Random {i}" for i in range(10)] + new_var_names = new_var_names_init + adata4.var_names[10:].to_list() + adata4.var_names = new_var_names + +def test_scvix_scarches_embedding(save_path): + # test transfer_anndata_setup + view + adata1 = synthetic_iid() + SCVIX.setup_anndata(adata1, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata1, batch_representation="embedding") + model.train(1, train_size=0.5) + dir_path = os.path.join(save_path, "saved_model/") + model.save(dir_path, overwrite=True) + + # adata2 has more genes and a perfect subset of adata1 + adata2 = synthetic_iid(n_genes=110) + adata2.obs["batch"] = adata2.obs.batch.cat.rename_categories(["batch_2", "batch_3"]) + SCVIX.prepare_query_anndata(adata2, dir_path) + SCVIX_query = SCVIX.load_query_data(adata2, dir_path) + SCVIX_query.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + adata3 = SCVIX.prepare_query_anndata(adata2, dir_path, inplace=False) + SCVIX_query2 = SCVIX.load_query_data(adata3, dir_path) + SCVIX_query2.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + # adata4 has more genes and missing 10 genes from adata1 + adata4 = synthetic_iid(n_genes=110) + new_var_names_init = [f"Random {i}" for i in range(10)] + new_var_names = new_var_names_init + adata4.var_names[10:].to_list() + adata4.var_names = new_var_names + +def test_scvix_minified(): + adata = synthetic_iid() + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") + model = SCVIX(adata, batch_representation="embedding", gene_likelihood="zinb") + model.train(1, train_size=0.5) + + qzm, qzv = model.get_latent_representation(give_mean=False, return_dist=True) + model.adata.obsm["X_latent_qzm"] = qzm + model.adata.obsm["X_latent_qzv"] = qzv + lib_size = np.squeeze(np.asarray(adata.X.sum(axis=-1))) + + scvi.settings.seed = 1 + params_orig = model.get_likelihood_parameters(n_samples=200, give_mean=True) + adata_orig = adata.copy() + + model.minify_adata() + assert model.minified_data_type == ADATA_MINIFY_TYPE.LATENT_POSTERIOR + assert model.adata_manager.registry is model.registry_ + + assert not _is_minified(adata) + assert adata is not model.adata + + orig_obs_df = adata_orig.obs + orig_obs_df[BaseMinifiedModeModelClass._OBSERVED_LIB_SIZE_KEY] = lib_size + assert model.adata.obs.equals(orig_obs_df) + assert model.adata.var_names.equals(adata_orig.var_names) + assert model.adata.var.equals(adata_orig.var) + + scvi.settings.seed = 1 + keys = ["mean", "dispersions", "dropout"] + params_latent = model.get_likelihood_parameters(n_samples=200, give_mean=True) + for k in keys: + assert params_latent[k].shape == params_orig[k].shape + + for k in keys: + assert_approx_equal(params_latent[k], params_orig[k]) From 4f1cde9aab06b0cf0389444dad32371daf795cf2 Mon Sep 17 00:00:00 2001 From: cane11 Date: Wed, 5 Nov 2025 14:00:01 -0800 Subject: [PATCH 11/24] Updates for CELLxGENE --- src/scvi/_constants.py | 1 + src/scvi/external/scvix/_model.py | 9 +++++++++ src/scvi/external/scvix/_module.py | 5 +++++ src/scvi/train/_trainingplans.py | 27 +++++++++++++++++++++------ 4 files changed, 36 insertions(+), 6 deletions(-) diff --git a/src/scvi/_constants.py b/src/scvi/_constants.py index d06c8aabcb..230a897037 100644 --- a/src/scvi/_constants.py +++ b/src/scvi/_constants.py @@ -7,6 +7,7 @@ class _REGISTRY_KEYS_NT(NamedTuple): BATCH_KEY: str = "batch" SITE_KEY: str = "site" ASSAY_KEY: str = "assay" + ADVERSARIAL_GROUP_KEY: str = "adversarial_group" SAMPLE_KEY: str = "sample" LABELS_KEY: str = "labels" PROTEIN_EXP_KEY: str = "proteins" diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index efca6dfc35..a4dd7106af 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -186,6 +186,7 @@ def __init__( n_input=self.summary_stats.n_vars, n_batch=self.summary_stats.n_batch, n_assay=self.summary_stats.n_assay, + n_adversarial_group=self.summary_stats.get("n_adversarial_group", 1), n_labels=self.summary_stats.get("n_labels", 1), n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), n_cats_per_cov=n_cats_per_cov, @@ -337,6 +338,7 @@ def setup_anndata( batch_key: str | None = None, assay_key: str | None = None, labels_key: str | None = None, + adversarial_group_key: str | None = None, unlabeled_category: str = "unlabeled", categorical_covariate_keys: list[str] | None = None, continuous_covariate_keys: list[str] | None = None, @@ -352,6 +354,9 @@ def setup_anndata( assay_key Key in ``adata.obs`` that corresponds to the assay of the data. %(param_labels_key)s + adversarial_group_key + Key in ``adata.obs`` that corresponds to the adversarial group for adversarial + training. If ``None``, performs no conditional adversarial training. %(param_unlabeled_category)s %(param_cat_cov_keys)s %(param_cont_cov_keys)s @@ -368,6 +373,10 @@ def setup_anndata( anndata_fields.append( LabelsWithUnlabeledObsField( REGISTRY_KEYS.LABELS_KEY, labels_key, unlabeled_category)) + if adversarial_group_key is not None: + anndata_fields.append( + CategoricalObsField( + REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, adversarial_group_key)) # register new fields if the adata is minified adata_minify_type = _get_adata_minify_type(adata) if adata_minify_type is not None: diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index 58f8c7c5ab..a70230820e 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -132,6 +132,7 @@ def __init__( n_batch: int = 0, n_assay: int = 0, n_labels: int = 0, + n_adversarial_group: int = 0, n_hidden: int = 128, n_latent: int = 10, n_layers: int = 1, @@ -168,6 +169,7 @@ def __init__( self.n_batch = n_batch self.n_assay = n_assay self.n_labels = n_labels + self.n_adversarial_group = n_adversarial_group self.encode_covariates = encode_covariates self.use_observed_lib_size = True self.n_hidden = n_hidden @@ -334,6 +336,7 @@ def _get_inference_input( MODULE_KEYS.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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.ADVERSARIAL_GROUP_KEY: tensors.get(REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, None), } else: return { @@ -365,6 +368,7 @@ def _regular_inference( assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, + adversarial_group: torch.Tensor | None = None, n_samples: int = 1, ) -> dict[str, torch.Tensor | Distribution | None]: """Run the regular inference process.""" @@ -406,6 +410,7 @@ def _regular_inference( MODULE_KEYS.Z_KEY: z, MODULE_KEYS.QZ_KEY: qz, MODULE_KEYS.LIBRARY_KEY: library, + MODULE_KEYS.ADVERSARIAL_GROUP_KEY: adversarial_group, } @auto_move_data diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index cff62bb66c..832e527797 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -655,7 +655,7 @@ def __init__( self.adversarial_classifier = False else: self.adversarial_classifier = Classifier( - n_input=self.module.n_latent, + n_input=self.module.n_latent+self.module.n_adversarial_group, n_hidden=128, n_labels=self.n_output_classifier, n_layers=2, @@ -663,14 +663,24 @@ def __init__( use_batch_norm=False, use_layer_norm=True, ) + else: self.adversarial_classifier = adversarial_classifier self.scale_adversarial_loss = scale_adversarial_loss self.automatic_optimization = False - def loss_adversarial_classifier(self, z, batch_index, predict_true_class=True): + def loss_adversarial_classifier(self, z, adversarial_group, batch_index, predict_true_class=True): """Loss for adversarial classifier.""" n_classes = self.n_output_classifier + adversarial_group_ = torch.nn.functional.one_hot( + adversarial_group, num_classes=self.module.n_adversarial_group + ).float() + if predict_true_class: # train classifier + z = z.detach() + #else: # fool classifier + # adversarial_group_emb = adversarial_group_emb.detach() + + z = torch.cat([z, adversarial_group_], dim=1) cls_logits = self.adversarial_classifier(z) if predict_true_class: @@ -706,11 +716,16 @@ def training_step(self, batch, batch_idx): inference_outputs, _, scvi_loss = self.forward(batch, loss_kwargs=self.loss_kwargs) z = inference_outputs["z"] + adversarial_group = inference_outputs.get("adversarial_group", None) + if adversarial_group is None: + adversarial_group = torch.zeros(z.size(0)).to(z.device).long() + else: + adversarial_group = adversarial_group.squeeze(-1).long() loss = scvi_loss.loss orig_loss = loss # fool classifier if doing adversarial training if kappa > 0 and self.adversarial_classifier is not False: - fool_loss = self.loss_adversarial_classifier(z, batch_tensor, False) + fool_loss = self.loss_adversarial_classifier(z, adversarial_group, batch_tensor, False) loss += fool_loss * kappa self.log("train_loss", loss, on_step=self.on_step, on_epoch=self.on_epoch, prog_bar=True) @@ -730,8 +745,8 @@ def training_step(self, batch, batch_idx): loss = 0. for i in range(self.adversarial_steps): qz = inference_outputs["qz"] - z = qz.sample().detach() - loss_ = kappa * self.loss_adversarial_classifier(z, batch_tensor, True) + z = qz.sample() + loss_ = kappa * self.loss_adversarial_classifier(z, adversarial_group, batch_tensor, True) if isinstance(self.module.prior, MogPrior) or isinstance(self.module.prior, VampPrior): qz_m, qz_v = qz.loc.detach(), qz.scale.detach() loss_ += self.module.prior.kl( @@ -794,7 +809,7 @@ def configure_optimizers(self): if self.adversarial_classifier is not False: params2 = filter(lambda p: p.requires_grad, self.adversarial_classifier.parameters()) optimizer2 = torch.optim.Adam( - params2, lr=1e-3, eps=0.01, weight_decay=self.weight_decay + params2, lr=3e-4, eps=1e-4, weight_decay=1e-9 ) config2 = {"optimizer": optimizer2} From e83be7f66ed2e864f6a860a3483f00130236bfc3 Mon Sep 17 00:00:00 2001 From: cane11 Date: Wed, 5 Nov 2025 14:00:24 -0800 Subject: [PATCH 12/24] Other changes --- src/scvi/module/_constants.py | 1 + src/scvi/train/_trainingplans.py | 3 --- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/src/scvi/module/_constants.py b/src/scvi/module/_constants.py index 4a75333db0..6626022bea 100644 --- a/src/scvi/module/_constants.py +++ b/src/scvi/module/_constants.py @@ -12,6 +12,7 @@ class _MODULE_KEYS(NamedTuple): QL_KEY: str = "ql" BATCH_INDEX_KEY: str = "batch_index" ASSAY_INDEX_KEY: str = "assay_index" + ADVERSARIAL_GROUP_KEY: str = "adversarial_group" SITE_INDEX_KEY: str = "site_index" Y_KEY: str = "y" CONT_COVS_KEY: str = "cont_covs" diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 832e527797..8fa45b2236 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -677,9 +677,6 @@ def loss_adversarial_classifier(self, z, adversarial_group, batch_index, predict ).float() if predict_true_class: # train classifier z = z.detach() - #else: # fool classifier - # adversarial_group_emb = adversarial_group_emb.detach() - z = torch.cat([z, adversarial_group_], dim=1) cls_logits = self.adversarial_classifier(z) From 15517ffeedc12d09d7ed94c6d1367b8e303f9226 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Wed, 5 Nov 2025 14:09:22 -0800 Subject: [PATCH 13/24] Conditional norm single class exception --- src/scvi/nn/_base_components.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index dd1a5573e6..dec01f161e 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -114,6 +114,10 @@ def __init__( self.n_continuous = n_continuous self.cond_cat = conditional_category + if conditional_norm and self.n_cat_list[self.cond_cat]==0: + raise ValueError( + "Conditional normalization is not applicable for a categorical variable with only one category." + ) self.n_cov = n_continuous + sum(self.n_cat_list) From 03a378218528ba01845bae2f66a30e857da82ac3 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Fri, 22 May 2026 11:15:58 +0200 Subject: [PATCH 14/24] removed muANVI --- src/scvi/external/__init__.py | 2 - src/scvi/external/muanvi/__init__.py | 4 - src/scvi/external/muanvi/_base_components.py | 238 ------ src/scvi/external/muanvi/_model.py | 725 ------------------- src/scvi/external/muanvi/_module.py | 477 ------------ src/scvi/external/muanvi/_utils.py | 182 ----- 6 files changed, 1628 deletions(-) delete mode 100644 src/scvi/external/muanvi/__init__.py delete mode 100644 src/scvi/external/muanvi/_base_components.py delete mode 100644 src/scvi/external/muanvi/_model.py delete mode 100644 src/scvi/external/muanvi/_module.py delete mode 100644 src/scvi/external/muanvi/_utils.py diff --git a/src/scvi/external/__init__.py b/src/scvi/external/__init__.py index bf642bcce0..dfa15451d9 100644 --- a/src/scvi/external/__init__.py +++ b/src/scvi/external/__init__.py @@ -11,7 +11,6 @@ from .methylvi import METHYLANVI, METHYLVI from .mrvi import MRVI from .mrvi_torch import TorchMRVI -from .muanvi import MUANVI from .poissonvi import POISSONVI from .resolvi import RESOLVI from .scar import SCAR @@ -47,7 +46,6 @@ "SCVIVA", "CYTOVI", "DIAGVI", - "MUANVI", ] diff --git a/src/scvi/external/muanvi/__init__.py b/src/scvi/external/muanvi/__init__.py deleted file mode 100644 index fde12a9392..0000000000 --- a/src/scvi/external/muanvi/__init__.py +++ /dev/null @@ -1,4 +0,0 @@ -from ._model import MUANVI -from ._module import MUANVAE - -__all__ = ["MUANVI", "MUANVAE"] diff --git a/src/scvi/external/muanvi/_base_components.py b/src/scvi/external/muanvi/_base_components.py deleted file mode 100644 index 485040a764..0000000000 --- a/src/scvi/external/muanvi/_base_components.py +++ /dev/null @@ -1,238 +0,0 @@ -import torch -from torch import nn as nn -from torch.nn import functional as F - -from scvi.module import Classifier -from scvi.nn import FCLayers - - -class Hierarchical_Classifier(nn.Module): - """ - Hierarchical Embedding Network - - Parameters (same as Classifier ) - ---------- - n_input - Number of input dimensions (dimensions of the latent space) - num_classes - number of labels in each label level in hierarchical list (ex : [2, 7]) - n_hidden - Number of hidden nodes in one layer - n_layers - Number of hidden layers per NN (per independent representation) - n_output - Number of dimensions of each independent representation - dropout_rate - dropout_rate for nodes - use_batch_norm - Whether to use batch norm in layers - use_layer_norm - Whether to use layer norm in layers - concatenation - Whether to concatenate or not the independent representations between layers - By default no concatenation. - """ - - def __init__( - self, - n_input: int, - num_classes: list, - n_hidden: int = 128, - dropout_rate: float = 0.1, - activation_fn: nn.Module = nn.ReLU, - n_layers: int = 3, - use_batch_norm: bool = False, - use_layer_norm: bool = True, - concatenation: bool = False, - logits: bool = True, # noqa, not used, we always return logits - ): - super().__init__() - self.n_input = n_input - self.n_hidden = n_hidden - self.concatenation = concatenation - - # independant representation level 1 of root level - layers = [ - FCLayers( - n_in=n_input, - n_out=n_hidden, - n_layers=n_layers, - n_hidden=n_hidden, - dropout_rate=dropout_rate, - use_batch_norm=use_batch_norm, - use_layer_norm=use_layer_norm, - activation_fn=activation_fn, - ) - for _ in range(len(num_classes)) - ] - # neural networks to obtain independant representations of dim n_output : - self.lvls = nn.ModuleList([nn.Sequential(layer) for layer in layers]) - - self.logits = nn.ModuleList([nn.Linear(n_hidden, num_class) for num_class in num_classes]) - self.softmax = nn.Softmax(dim=-1) - - def forward(self, x): - lvl_independents = [lvl(x) for lvl in self.lvls] - - logits_level = [ - logit(lvl_independent) - for logit, lvl_independent in zip(self.logits, lvl_independents, strict=False) - ] - probs_level = [self.softmax(logit_level) for logit_level in logits_level] - return probs_level, logits_level - - -class HierarchicalLossNetwork(Hierarchical_Classifier): - """ - Parameters (same as Classifier) - - ---------- - n_input - Number of input dimensions (dimensions of the latent space) - num_classes - number of labels in each class in hierarchical list (ex : [2, 7]) - n_hidden - Number of hidden nodes in one layer - n_layers - Number of hidden layers per NN (per independant representation) - n_output - Number of dimensions of each independant representation - dropout_rate - dropout_rate for nodes - use_batch_norm - Whether to use batch norm in layers - use_layer_norm - Whether to use layer norm in layers - activation_fn - Valid activation function from torch.nn - """ - - def __init__( - self, - n_input: int, - num_classes: list, - n_hidden: int = 128, - n_layers: int = 1, - dropout_rate: float = 0.1, - activation_fn: nn.Module = nn.ReLU, - use_batch_norm: bool = True, - use_layer_norm: bool = False, - **cls_parameters, - ): - # initialize the Classifier - super().__init__( - n_input=n_input, - n_hidden=n_hidden, - dropout_rate=dropout_rate, - activation_fn=activation_fn, - n_layers=n_layers, - use_batch_norm=use_batch_norm, - use_layer_norm=use_layer_norm, - num_classes=num_classes, - **cls_parameters, - ) - - self.total_level = len(num_classes) - - def calculate_lloss(self, predictions, true_labels, device, weights=None): - """ - Calculates the layer loss across all levels (multiple Cross-Entropy) - - Parameters - ---------- - predictions - Predictions of the model - true labels - Ground truth - weights - If we want to compute weighted cross entropy - """ - lloss = 0 - for l in range(self.total_level): - lloss += nn.CrossEntropyLoss()(predictions[l], true_labels[l]) - return lloss - - -class MultiBatchClassifier(nn.Module): - """ - Dictionary of batch-specific fully-connected NN classifiers. - - Parameters : parameters of the batch-specific classifiers - ---------- - n_input - Number of input dimensions - n_hidden - Number of nodes in hidden layer(s). If `0`, the classifier only consists of a - single linear layer. - n_labels - Numput of outputs dimensions - n_layers - Number of hidden layers. If `0`, the classifier only consists of a single - linear layer. - dropout_rate - dropout_rate for nodes - logits - Return logits or not - use_batch_norm - Whether to use batch norm in layers - use_layer_norm - Whether to use layer norm in layers - activation_fn - Valid activation function from torch.nn - n_sites - Number of different batch-specific classifiers to create - **kwargs - Keyword arguments passed into :class:`~scvi.nn.FCLayers`. - """ - - def __init__( - self, - n_input: int, - n_sites: int, - n_hidden: int = 128, - n_labels: int = 5, - n_layers: int = 1, - dropout_rate: float = 0.1, - use_batch_norm: bool = True, - use_layer_norm: bool = False, - activation_fn: nn.Module = nn.ReLU, - **kwargs, - ): - super().__init__() - self.n_sites = n_sites - self.classifier_dict = nn.ModuleDict( - { - str(i): Classifier( - n_input=n_input, - n_hidden=n_hidden, - n_labels=n_labels, - n_layers=n_layers, - dropout_rate=dropout_rate, - use_batch_norm=use_batch_norm, - use_layer_norm=use_layer_norm, - activation_fn=activation_fn, - **kwargs, - ) - for i in range(n_sites) - } - ) - - def forward(self, x, site_index): - """Forward computation for one mini batch of observations.""" - logits_list = [] - indices_list = [] - unique_sites = torch.unique(site_index) - - for site in unique_sites: - site_indices = torch.nonzero(site_index == site, as_tuple=True)[0].to(x.device) - indices_list.append(site_indices) - x_site = x[site_indices] - logits_list.append(self.classifier_dict[str(int(site.item()))](x_site)) - all_logits = torch.cat(logits_list, dim=0).to(x.device) - all_indices = torch.cat(indices_list, dim=0) - - # Sort the indices to get the original minibatch order - sorted_indices = torch.argsort(all_indices).to(x.device) - output_logits = all_logits[sorted_indices] - - return F.softmax(output_logits, dim=-1), output_logits \ No newline at end of file diff --git a/src/scvi/external/muanvi/_model.py b/src/scvi/external/muanvi/_model.py deleted file mode 100644 index 8daa3e0c21..0000000000 --- a/src/scvi/external/muanvi/_model.py +++ /dev/null @@ -1,725 +0,0 @@ -import logging -import warnings -from collections.abc import Sequence -from copy import deepcopy -from typing import Literal - -import numpy as np -import pandas as pd -import torch -from anndata import AnnData - -from scvi import REGISTRY_KEYS -from scvi.data import AnnDataManager -from scvi.data._constants import _SETUP_ARGS_KEY -from scvi.data.fields import ( - CategoricalJointObsField, - CategoricalObsField, - LabelsWithUnlabeledObsField, - LayerField, - NumericalJointObsField, - NumericalObsField, -) -from scvi.dataloaders import SemiSupervisedDataSplitter -from scvi.model import SCVI -from scvi.model._utils import get_max_epochs_heuristic, parse_device_args -from scvi.model.base import ArchesMixin, BaseModelClass, RNASeqMixin, VAEMixin -from scvi.model.base._archesmixin import _get_loaded_data, _set_params_online_update -from scvi.model.base._save_load import ( - _initialize_model, - _validate_var_names, -) -from scvi.train import SemiSupervisedTrainingPlan, TrainRunner -from scvi.train._callbacks import SubSampleLabels -from scvi.utils import setup_anndata_dsp -from scvi.utils._docstrings import devices_dsp - -from ._module import MUANVAE -from ._utils import LabelsWithUnlabeledJointObsField, _get_site_code_from_category - -logger = logging.getLogger(__name__) - - -class MUANVI(RNASeqMixin, VAEMixin, ArchesMixin, BaseModelClass): - """ - Hierarchical multi-annotator Variational Inference [Xu21]_. - - Inspired from M1 + M2 model, as described in (https://arxiv.org/pdf/1406.5298.pdf). - - Parameters - ---------- - adata - AnnData object that has been registered via :meth:`~scvi.model.MUANVI.setup_anndata`. - n_hidden - Number of nodes per hidden layer. - n_latent - Dimensionality of the latent space. - n_layers - Number of hidden layers used for encoder and decoder 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 - gene_likelihood - One of: - * ``'nb'`` - Negative binomial distribution - * ``'zinb'`` - Zero-inflated negative binomial distribution - * ``'poisson'`` - Poisson distribution - update_yprior - Whether to perform the hierarchical update of the y prior parameter in the loss - batches_to_harmonize - List of two indices of the two batches to label-harmonize. Has to be defined if the dataset has more than 2 batches. - **model_kwargs - Keyword args for :class:`~scvi.module.MUANVAE` - - Examples - -------- - >>> adata = anndata.read_h5ad(path_to_anndata) - >>> scvi.external.MUANVI.setup_anndata(adata, labels=["labels_0", "labels_1"], unknown_categories=["unknown, "unknown"]) - >>> model = scvi.external.MUANVI(adata) - >>> model.train() - >>> adata.obsm["X_muanvi"] = model.get_latent_representation() - >>> adata.obs["pred_label_coarse"] = model.predict()[0] - >>> adata.obs["pred_label_fine"] = model.predict()[1] - - """ - - _module_cls = MUANVAE - _training_plan_cls = SemiSupervisedTrainingPlan - - def __init__( - self, - adata: AnnData, - n_hidden: int = 128, - n_latent: int = 10, - n_layers: int = 1, - dropout_rate: float = 0.1, - dispersion: Literal["gene", "gene-batch", "gene-label", "gene-cell"] = "gene", - gene_likelihood: Literal["zinb", "nb", "poisson"] = "nb", - update_yprior: bool = True, - eps_yprior: float = 1e-4, - **model_kwargs, - ): - super().__init__(adata) - muanvae_model_kwargs = dict(model_kwargs) - self._set_indices_and_labels() - - n_batch = self.summary_stats.n_batch - n_fine_labels = self.summary_stats.n_labels - 1 - n_site = self.summary_stats.n_site - n_assay = self.summary_stats.n_assay - - self.hierarchy_dict, self.num_classes, self.hierarchy_matrix = self.extract_hierarchy( - n_site=n_site, eps_yprior=eps_yprior - ) - n_cats_per_cov = ( - self.adata_manager.get_state_registry(REGISTRY_KEYS.CAT_COVS_KEY).n_cats_per_key - if REGISTRY_KEYS.CAT_COVS_KEY in self.adata_manager.data_registry - else None - ) - - use_size_factor_key = REGISTRY_KEYS.SIZE_FACTOR_KEY in self.adata_manager.data_registry - - self.module = self._module_cls( - n_input=self.summary_stats.n_vars, - n_batch=n_batch, - n_site=n_site, - n_assay=n_assay, - n_fine_labels = n_fine_labels, - num_classes=self.num_classes, - n_continuous_cov=self.summary_stats.get("n_extra_continuous_covs", 0), - n_cats_per_cov=n_cats_per_cov, - n_hidden=n_hidden, - n_latent=n_latent, - n_layers=n_layers, - dropout_rate=dropout_rate, - dispersion=dispersion, - gene_likelihood=gene_likelihood, - use_size_factor_key=use_size_factor_key, - hierarchy_dict=self.hierarchy_dict, - hierarchy_matrix=self.hierarchy_matrix, - update_yprior=update_yprior, - **muanvae_model_kwargs, - ) - - self.unsupervised_history_ = None - self.semisupervised_history_ = None - - self._model_summary_string = ( - f"muANVI Model with the following params: \nunlabeled_category: {self.unlabeled_category}, n_hidden: {n_hidden}, n_latent: {n_latent}" - f", n_layers: {n_layers}, dropout_rate: {dropout_rate}, dispersion: {dispersion}, gene_likelihood: {gene_likelihood}" - ) - self.init_params_ = self._get_init_params(locals()) - self.was_pretrained = False - self.n_fine_labels = n_fine_labels - - @classmethod - def from_scvi_model( - cls, - scvi_model: SCVI, - unlabeled_category: list[str], - fine_labels: str | None = None, - label_hierarchy: list[str] | None = None, - adata: AnnData | None = None, - **muanvi_kwargs, - ): - """ - Initialize scHANVI model with weights from pretrained :class:`~scvi.model.SCVI` model. - - Parameters - ---------- - scvi_model - Pretrained scvi model - fine_labels - key in `adata.obs` for label information. If this value is not None, the key will - overwrite the `labels_key` used to setup AnnData with scvi. - label_hierarchy - List of strings, with levels of the cell-type hierarchy. Full hierarchy is inferred by - concatenating label_hierarchy and fine_labels. - unlabeled_category - Value used for unlabeled cells in `labels_key`. - adata - AnnData object that has been registered via :meth:`~scvi.model.MUANVI.setup_anndata`. - muanvi_kwargs - kwargs for muANVI model - """ - scvi_model._check_if_trained(message="Passed in scvi model hasn't been trained yet.") - - muanvi_kwargs = dict(muanvi_kwargs) - init_params = scvi_model.init_params_ - non_kwargs = init_params["non_kwargs"] - kwargs = init_params["kwargs"] - kwargs = {k: v for (i, j) in kwargs.items() for (k, v) in j.items()} - for k, v in {**non_kwargs, **kwargs}.items(): - if k in muanvi_kwargs.keys(): - warnings.warn( - f"Ignoring param '{k}' as it was already passed in to " - + f"pretrained scvi model with value {v}.", - stacklevel=2, - ) - del muanvi_kwargs[k] - - if adata is None: - adata = scvi_model.adata - else: - # validate new anndata against old model - scvi_model._validate_anndata(adata) - - scvi_setup_args = deepcopy(scvi_model.adata_manager.registry[_SETUP_ARGS_KEY]) - scvi_labels_key = scvi_setup_args["labels_key"] - if fine_labels is None and scvi_labels_key is None: - raise ValueError( - "A `labels_key` list is necessary as the SCVI model was initialized without one." - ) - if fine_labels is not None: - scvi_setup_args.update({"fine_labels": fine_labels}) - scvi_setup_args.pop("labels_key", None) - else: - scvi_setup_args["fine_labels"] = scvi_setup_args.pop("labels_key") - - cls.setup_anndata( - adata, - unlabeled_category=unlabeled_category, - label_hierarchy=label_hierarchy, - **scvi_setup_args, - ) - muanvi_model = cls(adata, **non_kwargs, **kwargs, **muanvi_kwargs) - scvi_state_dict = scvi_model.module.state_dict() - muanvi_model.module.load_state_dict(scvi_state_dict, strict=False) - muanvi_model.was_pretrained = True - - return muanvi_model - - def _set_indices_and_labels(self): - """Set indices for labeled and unlabeled cells.""" - labels_state_registry = self.adata_manager.get_state_registry("label_hierarchy") - self.original_label_keys = labels_state_registry.field_keys - self.unlabeled_category = labels_state_registry.unlabeled_category - - # Dataframe of 2 columns for the 2 layers of labels - labels = {field: self.adata.obs[field] for field in self.original_label_keys} - self.labels = pd.DataFrame(labels) - self._label_mapping = labels_state_registry.mappings.to_dict() - # a cell is unlabeled if it is not labeled at the finest state - labeled_indices_list = [ - set(np.where(self.labels.iloc[:, idx] != self.unlabeled_category)[0]) - for idx, _ in enumerate(self.original_label_keys) - ] - self._labeled_indices = list(set.intersection(*labeled_indices_list)) - self._unlabeled_indices = list( - set(np.arange(self.adata.n_obs)) - set(self._labeled_indices) - ) - self._code_to_label = { - layer: dict(enumerate(self._label_mapping[layer])) for layer in self._label_mapping - } - - def extract_hierarchy(self, n_site, eps_yprior): - """ - Method to extract automatically the intrinsic hierarchy in the data. - - Parameters - ---------- - n_site - Number of different vocabulary sites in the data - eps_yprior - Epsilon value to add to the y prior parameter to avoid overconfidence - """ - labels_state_registry = self.adata_manager.get_state_registry("label_hierarchy") - label_keys = labels_state_registry.field_keys - - def fixed_depth_groupby(df, label_keys): - if len(label_keys)==1: - return list(df[label_keys[-1]].unique()) if not df.empty else [] - else: - # Recursive case: build dictionaries up to the fixed depth - return { - key: fixed_depth_groupby(sub_df, label_keys[1:]) - for key, sub_df in df.groupby(label_keys[0], observed=True) - } - - num_classes = labels_state_registry.n_cats_per_key - hierarchy_dict = fixed_depth_groupby(self.labels, label_keys) - hierarchy_matrix = [] - - for n_label in range(1, len(num_classes)): - curr = pd.DataFrame( - 0, - index=np.arange(num_classes[n_label - 1]), - columns=np.arange(num_classes[n_label]), - ) - # Site specific last layer. - if n_label == len(num_classes) - 1: - hierarchy_matrix_ = torch.zeros( - num_classes[n_label - 1], num_classes[n_label], n_site - ) - for site in range(n_site): - adata = self.adata[self.adata.obs["_scvi_site"] == site] - - if adata.n_obs > 0: - curr_ = pd.crosstab( - adata.obsm["_scvi_label_hierarchy"][ - label_keys[n_label - 1] - ], - adata.obsm["_scvi_label_hierarchy"][label_keys[n_label]], - ) - curr_[curr_ > 0] = 1 - curr_ = curr_.loc[curr.index, curr.columns] - curr = curr_.div(curr_.sum(axis=1), axis=0).fillna(0) - hierarchy_matrix_[:, :, site] = torch.tensor(curr.values) - else: - curr_ = pd.crosstab( - self.adata.obsm["_scvi_label_hierarchy"][label_keys[n_label - 1]], - self.adata.obsm["_scvi_label_hierarchy"][label_keys[n_label]], - ) - curr_[curr_ > 0] = 1 - curr_ = curr_.loc[curr.index, curr.columns] - curr = curr_.div(curr_.sum(axis=1), axis=0).fillna(0) - hierarchy_matrix_ = torch.tensor(curr.values) - - hierarchy_matrix.append(hierarchy_matrix_) - - return (hierarchy_dict, num_classes, hierarchy_matrix) - - def predict( - self, - adata: AnnData | None = None, - indices: Sequence[int] | None = None, - soft: bool = False, - batch_size: int | None = None, - level: int = -1, - sites_to_predict: int | str | str | None = None, - ) -> np.ndarray | pd.DataFrame: - """ - Return cell label predictions. - - Parameters - ---------- - adata - AnnData object that has been registered via :meth:`~scvi.model.SCANVI.setup_anndata`. - indices - indices for which to return probabilities. - soft - If True, returns per class probabilities - batch_size - Minibatch size for data loading into model. Defaults to `scvi.settings.batch_size`. - site_to_predict - For cross prediction purposes : sites to use when cross-predicting labels. - If None, normal prediction occurs and each cell is labeled accordingly to - its own site-specific classifier. Can be list of strings to use multiple sites, - or a string to use a single site. - level - Level of the hierarchy to predict. If -1, predicts at the finest level. - """ - adata = self._validate_anndata(adata) - - if indices is None: - indices = np.arange(adata.n_obs) - - scdl = self._make_data_loader( - adata=adata, - indices=indices, - batch_size=batch_size, - ) - # total depth of the hierarchy - total_level = len(self.module.num_classes) - class_labels = self.adata_manager.get_state_registry("label_hierarchy").field_keys - - sites_to_predict_, site_mappings_ = _get_site_code_from_category( - self.get_anndata_manager(adata, required=True), sites_to_predict - ) - pred = {str(site_mappings_[i]) + "_fine": [] for i in sites_to_predict_ if i is not None} - if None in sites_to_predict_: - for i in class_labels: - pred[i] = [] - - for _, tensors in enumerate(scdl): - for site_to_predict_ in sites_to_predict_: - probs_, _ = self.module.classification(tensors, site_to_predict=site_to_predict_) - if site_to_predict_ is not None: - if not soft: - pred_ = probs_.argmax(dim=1) - else: - pred_ = probs_ - pred[site_mappings_[site_to_predict_] + "_fine"].append(pred_.detach().cpu()) - - else: - # pred is a tuple of probabilities and logits for each layer - probs = probs_ - - for i in range(total_level): - if not soft: - pred[class_labels[i]].append(probs[i].argmax(dim=1).cpu()) - else: - pred[class_labels[i]].append(probs[i].detach().cpu()) - - for key in pred.keys(): - pred[key] = torch.cat(pred[key]).numpy() - if not soft: - pred[key] = [self._code_to_label[key][ct] for ct in pred[key]] - - if not soft: - pred = pd.DataFrame.from_dict(pred) - pred.index = adata.obs_names[indices] - return pred - else: - for key in pred.keys(): - columns = list(self._code_to_label[key].values())[:-1] - - pred[key] = pd.DataFrame( - pred[key], - columns=columns, - index=adata.obs_names[indices], - ) - return pred - - @classmethod - @devices_dsp.dedent - def load_query_data( - cls, - adata: AnnData, - reference_model: str | BaseModelClass, - inplace_subset_query_vars: bool = False, - accelerator: str = "auto", - device: int | str = "auto", - unfrozen: bool = False, - freeze_dropout: bool = False, - freeze_expression: bool = True, - freeze_decoder_first_layer: bool = True, - freeze_batchnorm_encoder: bool = True, - freeze_batchnorm_decoder: bool = False, - freeze_classifier: bool = True, - ): - """Online update of a reference model with scArches algorithm :cite:p:`Lotfollahi21`. - - Parameters - ---------- - adata - AnnData organized in the same way as data used to train model. - It is not necessary to run setup_anndata, - as AnnData is validated against the ``registry``. - reference_model - Either an already instantiated model of the same class, or a path to - saved outputs for reference model. - inplace_subset_query_vars - Whether to subset and rearrange query vars inplace based on vars used to - train reference model. - %(param_accelerator)s - %(param_device)s - unfrozen - Override all other freeze options for a fully unfrozen model - freeze_dropout - Whether to freeze dropout during training - freeze_expression - Freeze neurons corersponding to expression in first layer - freeze_decoder_first_layer - Freeze neurons corresponding to first layer in decoder - freeze_batchnorm_encoder - Whether to freeze batchnorm weight and bias during training for encoder - freeze_batchnorm_decoder - Whether to freeze batchnorm weight and bias during training for decoder - freeze_classifier - Whether to freeze classifier completely. Only applies to `SCANVI`. - """ - _, _, device = parse_device_args( - accelerator=accelerator, - devices=device, - return_device="torch", - validate_single_device=True, - ) - - attr_dict, var_names, load_state_dict = _get_loaded_data(reference_model, device=device) - - if inplace_subset_query_vars: - logger.debug("Subsetting query vars to reference vars.") - adata._inplace_subset_var(var_names) - _validate_var_names(adata, var_names) - - registry = attr_dict.pop("registry_") - if _SETUP_ARGS_KEY not in registry: - raise ValueError( - "Saved model does not contain original setup inputs. " - "Cannot load the original setup." - ) - - cls.setup_anndata( - adata, - source_registry=registry, - extend_categories=True, - allow_missing_labels=True, - **registry[_SETUP_ARGS_KEY], - ) - - model = _initialize_model(cls, adata, attr_dict) - adata_manager = model.get_anndata_manager(adata, required=True) - - if REGISTRY_KEYS.CAT_COVS_KEY in adata_manager.data_registry: - raise NotImplementedError( - "scArches currently does not support models with extra categorical covariates." - ) - - model.to_device(device) - - # model tweaking - new_state_dict = model.module.state_dict() - additional_parameters = set() - for key, new_ten in new_state_dict.items(): - load_ten = load_state_dict.get(key, None) - if load_ten is None: - # Picks up that additional site classifier was added, makes it trainable by default. - if "y_prior_fine" not in key: - additional_parameters.add(key) # TODO check this. - load_state_dict[key] = new_ten - continue - if new_ten.size() == load_ten.size(): - continue - # new categoricals changed size - else: - if new_ten.size()[0] != load_ten.size()[0]: - new_ten = new_ten.to(load_ten.device) - dim_diff = new_ten.size()[0] - load_ten.size()[0] - fixed_ten = torch.cat([load_ten, new_ten[-dim_diff:, ...]], dim=0) - load_state_dict[key] = fixed_ten - else: - new_ten = new_ten.to(load_ten.device) - dim_diff = new_ten.size()[-1] - load_ten.size()[-1] - fixed_ten = torch.cat([load_ten, new_ten[..., -dim_diff:]], dim=-1) - load_state_dict[key] = fixed_ten - - model.module.load_state_dict(load_state_dict) - model.module.eval() - - _set_params_online_update( - model.module, - unfrozen=unfrozen, - freeze_decoder_first_layer=freeze_decoder_first_layer, - freeze_batchnorm_encoder=freeze_batchnorm_encoder, - freeze_batchnorm_decoder=freeze_batchnorm_decoder, - freeze_dropout=freeze_dropout, - freeze_expression=freeze_expression, - freeze_classifier=freeze_classifier, - parameters_yes_grad=additional_parameters, - ) - model.is_trained_ = False - - return model - - def train( - self, - max_epochs: int | None = None, - n_samples_per_label: float | None = None, - check_val_every_n_epoch: int | None = None, - train_size: float = 0.9, - validation_size: float | None = None, - shuffle_set_split: bool = True, - batch_size: int = 128, - accelerator: str = "auto", - devices: int | list[int] | str = "auto", - datasplitter_kwargs: dict | None = None, - plan_kwargs: dict | None = None, - **trainer_kwargs, - ): - """Train the model. - - Parameters - ---------- - max_epochs - Number of passes through the dataset for semisupervised training. - n_samples_per_label - Number of subsamples for each label class to sample per epoch. By default, there - is no label subsampling. - check_val_every_n_epoch - Frequency with which metrics are computed on the data for validation set for both - the unsupervised and semisupervised trainers. If you'd like a different frequency for - the semisupervised trainer, set check_val_every_n_epoch in semisupervised_train_kwargs. - train_size - Size of training set in the range [0.0, 1.0]. - validation_size - Size of the test set. If `None`, defaults to 1 - `train_size`. If - `train_size + validation_size < 1`, the remaining cells belong to a test set. - shuffle_set_split - Whether to shuffle indices before splitting. If `False`, the val, train, and test set - are split in the sequential order of the data according to `validation_size` and - `train_size` percentages. - batch_size - Minibatch size to use during training. - %(param_accelerator)s - %(param_devices)s - datasplitter_kwargs - Additional keyword arguments passed into - :class:`~scvi.dataloaders.SemiSupervisedDataSplitter`. - plan_kwargs - Keyword args for :class:`~scvi.train.SemiSupervisedTrainingPlan`. Keyword arguments - passed to `train()` will overwrite values present in `plan_kwargs`, when appropriate. - **trainer_kwargs - Other keyword args for :class:`~scvi.train.Trainer`. - """ - if max_epochs is None: - max_epochs = get_max_epochs_heuristic(self.adata.n_obs) - - if self.was_pretrained: - max_epochs = int(np.min([10, np.max([2, round(max_epochs / 3.0)])])) - - plan_kwargs = {} if plan_kwargs is None else plan_kwargs - datasplitter_kwargs = datasplitter_kwargs or {} - - # if we have labeled cells, we want to subsample labels each epoch - sampler_callback = [SubSampleLabels()] if len(self._labeled_indices) != 0 else [] - - data_splitter = SemiSupervisedDataSplitter( - adata_manager=self.adata_manager, - train_size=train_size, - validation_size=validation_size, - shuffle_set_split=shuffle_set_split, - n_samples_per_label=n_samples_per_label, - batch_size=batch_size, - **datasplitter_kwargs, - ) - - warmup_epochs = plan_kwargs.pop("warmup_epochs", None) - - if warmup_epochs is not None and warmup_epochs > 0: - logger.info(f"Pretraining for {max_epochs} epochs.") - - plan_kwargs_pre = plan_kwargs.copy() - plan_kwargs_pre["warmup_model"] = True - plan_kwargs_pre["n_epochs_kl_warmup"] = warmup_epochs - - training_plan = self._training_plan_cls( - self.module, n_classes=self.n_fine_labels, **plan_kwargs_pre - ) # n_classes set at the finest level to track accuracy at that level. - runner_pre = TrainRunner( - self, - training_plan=training_plan, - data_splitter=data_splitter, - max_epochs=warmup_epochs, - accelerator=accelerator, - devices=devices, - check_val_every_n_epoch=check_val_every_n_epoch, - **trainer_kwargs, - ) - runner_pre() - self.was_pretrained = True - - logger.info(f"Training for {max_epochs} epochs.") - - if "callbacks" in trainer_kwargs.keys(): - trainer_kwargs["callbacks"] + [sampler_callback] - else: - trainer_kwargs["callbacks"] = sampler_callback - training_plan = self._training_plan_cls( - self.module, n_classes=self.n_fine_labels, **plan_kwargs - ) - - runner = TrainRunner( - self, - training_plan=training_plan, - data_splitter=data_splitter, - max_epochs=max_epochs, - accelerator=accelerator, - devices=devices, - check_val_every_n_epoch=check_val_every_n_epoch, - **trainer_kwargs, - ) - - return runner() - - @classmethod - @setup_anndata_dsp.dedent - def setup_anndata( - cls, - adata: AnnData, - fine_labels_key: str, - unlabeled_category: list[str | int | float], - layer: str | None = None, - site_key: str | None = None, - assay_key: str | None = None, - batch_key: str | None = None, - size_factor_key: str | None = None, - categorical_covariate_keys: list[str] | None = None, - continuous_covariate_keys: list[str] | None = None, - label_hierarchy: list[str] | None = None, - **kwargs, - ): - """ - %(summary)s. - - Parameters - ---------- - %(param_layer)s - %(param_batch_key)s - %(param_site_key)s - %(param_assay_key)s - fine_labels_key - key in `adata.obs` for fine label information. Categories will automatically be - converted into integer categories and saved to `adata.obs['_scvi_labels']`. - If `None`, assigns the same label to all the data. This information can be - site-specific. In this case we expect all the labels in a single obs column. - This is analogous to the label key in scANVI. - %(param_size_factor_key)s - %(param_cat_cov_keys)s - %(param_cont_cov_keys)s - label_hierarchy - List of strings, with levels of the cell-type hierarchy. - The first list is the root level, the last list is the second finest level. - The full hierarchy is inferred by concatenating label_hierarchy and fine_labels. - """ - setup_method_args = cls._get_setup_method_args(**locals()) - anndata_fields = [ - LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), - CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), - CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), - CategoricalObsField(REGISTRY_KEYS.SITE_KEY, site_key), - LabelsWithUnlabeledObsField(REGISTRY_KEYS.LABELS_KEY, fine_labels_key, unlabeled_category), - NumericalObsField(REGISTRY_KEYS.SIZE_FACTOR_KEY, size_factor_key, required=False), - CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), - NumericalJointObsField(REGISTRY_KEYS.CONT_COVS_KEY, continuous_covariate_keys), - LabelsWithUnlabeledJointObsField( - "label_hierarchy", label_hierarchy+[fine_labels_key], unlabeled_category), - ] - adata_manager = AnnDataManager(fields=anndata_fields, setup_method_args=setup_method_args) - adata_manager.register_fields(adata, **kwargs) - cls.register_manager(adata_manager) diff --git a/src/scvi/external/muanvi/_module.py b/src/scvi/external/muanvi/_module.py deleted file mode 100644 index 688cd6ed0d..0000000000 --- a/src/scvi/external/muanvi/_module.py +++ /dev/null @@ -1,477 +0,0 @@ -from typing import Literal - -import torch -from torch.distributions import Categorical, Independent, MixtureSameFamily, Normal -from torch.distributions import kl_divergence as kl -from torch.nn import functional as F - -from scvi import REGISTRY_KEYS -from scvi.module import SCANVAE -from scvi.module._utils import broadcast_labels -from scvi.module.base import LossOutput, auto_move_data - -from ._base_components import HierarchicalLossNetwork, MultiBatchClassifier - - -class MUANVAE(SCANVAE): - """ - Single-cell multiple-annotation using variational inference. - - This is a re-implementation of a hierarchical cell-type annotation model - inspired from scANVI model described in [Xu21]_,. - - Parameters - ---------- - n_input - Number of input genes - n_batch - Number of batches - n_site - Number of annotation sites - n_assay - Number of assays - n_fine_labels - Number of fine labels - num_classes - Number of labels per class organized in a hierarchical list - n_hidden - Number of nodes per hidden layer - n_latent - Dimensionality of the latent space - n_layers - Number of hidden layers used for encoder and decoder NNs - n_continuous_cov - Number of continuous covariates - n_cats_per_cov - Number of categories for each extra categorical covariate - 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 - conditioning_class - index of the class conditioning the second latent space (0 being the coarse class, 1 being the fine class). Default : coarse labels. - hierarchy_matrix - Matrix representing the hierarchy, computed by scATVI. If None, no hierarchical y prior update in the loss. - use_batch_norm - Whether to use batch norm in layers - use_layer_norm - Whether to use layer norm in layers - prior_z1 - Whether to use MoG or simple Gaussian for prior of z1 - **vae_kwargs - Keyword args for :class:`~scvi.module.VAE` - """ - - def __init__( - self, - n_input: int, - num_classes: list, - hierarchy_dict: dict, - n_batch: int = 0, - n_site: int = 0, - n_assay: int = 0, - n_fine_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: str = "gene", - log_variational: bool = True, - gene_likelihood: str = "nb", - classifier_parameters: dict | None= None, - classifier_parameters_muanvae: dict | None= None, - use_batch_norm: Literal["encoder", "decoder", "none", "both"] = "none", - use_layer_norm: Literal["encoder", "decoder", "none", "both"] = "both", - conditioning_class: int = -1, - mog_class: int = 0, - hierarchy_matrix=None, - update_yprior=True, - prior_z1: str = "gaussian", - eps_yprior: float = 1e-6, - **scanvae_kwargs, - ): - self.conditioning_class = conditioning_class - self.mog_class = mog_class - self.site_specific_classifier = n_site > 1 - self.n_site = n_site - self.n_assay = n_assay - self.num_classes = num_classes - self.update_yprior = update_yprior - self.n_labels_conditioning = num_classes[self.conditioning_class] - self.n_fine_labels = n_fine_labels - self.hiearchy_dict = hierarchy_dict - - if classifier_parameters is None: - classifier_parameters = {} - if classifier_parameters_muanvae is None: - classifier_parameters_muanvae = {} - - cls_parameters = { - "n_layers": n_layers, - "n_hidden": n_hidden, - "dropout_rate": dropout_rate, - "logits": True, - } - cls_parameters.update(classifier_parameters) - - super().__init__( - n_input, - n_batch=n_batch, - n_labels=self.n_fine_labels, - n_hidden=n_hidden, - n_latent=n_latent, - n_layers=n_layers, - n_continuous_cov=n_continuous_cov, - n_cats_per_cov=n_cats_per_cov, - dropout_rate=dropout_rate, - dispersion=dispersion, - log_variational=log_variational, - gene_likelihood=gene_likelihood, - classifier_parameters=classifier_parameters, - use_batch_norm=use_batch_norm, - use_layer_norm=use_layer_norm, - **scanvae_kwargs, - ) - cls_parameters.update(classifier_parameters_muanvae) - self.num_classes = num_classes - self.total_level = len(self.num_classes) # depth of hierarchy - self.prior_z1 = prior_z1 - - self.classifier = HierarchicalLossNetwork( - n_input=self.n_latent, - num_classes=self.num_classes[:-1], - **cls_parameters, - ) - # the site-specific classifiers must have layer norm (batch norm does not work if there is 1 single observation from a batch in a minibatch) - self.multi_classifier_fine = MultiBatchClassifier( - n_input=self.n_latent, - n_sites=n_site, - n_labels=n_fine_labels, - use_batch_norm=False, - use_layer_norm=True, - **cls_parameters, - ) - - # register y_prior on the fine labels - if not self.update_yprior: - hierarchy_matrix = [ - torch.tensor(hierarchy_matrix[i].sum(0) > 0, dtype=torch.float) - for i in range(len(num_classes) - 1) - ] - for i in range(0, self.total_level): - if i == self.total_level - 1: - self.y_prior_fine = torch.nn.ParameterList( - [ - torch.nn.Parameter( # - hierarchy_matrix[i - 1][:, :, site] + eps_yprior, - requires_grad=False, - ) - for site in range(n_site) - ] - ) - elif i == 0: - self.register_buffer( - f"y_prior_{i}", - torch.nn.Parameter( # - torch.full([num_classes[0]], 1 / num_classes[0], dtype=torch.float), - requires_grad=False, - ), - ) - else: - self.register_buffer( - f"y_prior_{i}", - torch.nn.Parameter( # - torch.tensor(hierarchy_matrix[i - 1] + eps_yprior, dtype=torch.float), - requires_grad=False, - ), - ) - if self.prior_z1 == "mog": - self.register_parameter( - "prior_z1_means", - torch.nn.Parameter(torch.randn([num_classes[self.mog_class], n_latent])), - ) - self.register_parameter( - "prior_z1_scales", - torch.nn.Parameter(torch.zeros([num_classes[self.mog_class], n_latent])), - ) - self.register_parameter( - "prior_z1_logits", torch.nn.Parameter(torch.ones([num_classes[self.mog_class]])) - ) - if self.prior_z1 == "mog_celltype": - self.register_parameter( - "prior_z1_means", - torch.nn.Parameter(torch.zeros([1, num_classes[self.mog_class], n_latent])), - ) - self.register_parameter( - "prior_z1_scales", - torch.nn.Parameter(torch.zeros([1, num_classes[self.mog_class], n_latent])), - ) - - @auto_move_data - def classify( - self, - x, - batch_index=None, - site_index=None, - cont_covs=None, - cat_covs=None, - site_to_predict: int | None = None, - precomputed_z: torch.Tensor | None = None, - ): - """ - Classify cells using the model. - - Parameters - ---------- - site_to_predict - For cross prediction purposes : index of the fine classifier to use when cross-predicting labels. - If None, normal prediction occurs, which is the case during training. - If not None, cross classification on only fine layer with this batch-specific classifier. - precomputed_z - Precomputed z1 latent space. If None, z1 is computed from x. - """ - if precomputed_z is not None: - z = precomputed_z - else: - if self.log_variational: # for numerical stability - x = torch.log(1 + 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 = () - qz, _ = self.z_encoder( - encoder_input, batch_index, *categorical_input - ) # q(z1|x) without the var qz_v - z = qz.rsample() - if site_to_predict is not None: - site_to_predict_index = torch.full((z.shape[0], 1), site_to_predict, dtype=torch.int32) - probs_fine, _ = self.multi_classifier_fine(z, site_to_predict_index) - return probs_fine, _ - probs_classifier, logits_classifier = self.classifier(z) - probs_fine, logits_fine = self.multi_classifier_fine(z, site_index) - probs_classifier += [probs_fine] - logits_classifier += [logits_fine] - - return probs_classifier, logits_classifier - - @auto_move_data - def classification( - self, - tensors, - return_classifier_loss=False, - site_to_predict=None, - precomputed_z=None, - ): - x = tensors[REGISTRY_KEYS.X_KEY] - y = tensors["label_hierarchy"] - batch_idx = tensors[REGISTRY_KEYS.BATCH_KEY] - site_idx = tensors[REGISTRY_KEYS.SITE_KEY] - cont_covs = tensors.get(REGISTRY_KEYS.CONT_COVS_KEY, None) - cat_covs = tensors.get(REGISTRY_KEYS.CAT_COVS_KEY, None) - - probs, logits = self.classify( - x, - batch_index=batch_idx, - site_index=site_idx, - cat_covs=cat_covs, - cont_covs=cont_covs, - site_to_predict=site_to_predict, - precomputed_z=precomputed_z, - ) - if not return_classifier_loss: - return probs, logits - - classification_loss = 0 - for l in range(self.total_level): - labels_curr_level = y[:, l].view(-1).long() - classification_loss += F.cross_entropy( - logits[l], - labels_curr_level, - ignore_index=self.num_classes[l], - reduction="mean" - ) - true_labels_fine = torch.unsqueeze(labels_curr_level, 1) - return classification_loss, true_labels_fine, logits[-1] - - def loss( - self, - tensors, - inference_outputs, - generative_ouputs, - kl_weight=1, - labelled_tensors=None, - classification_ratio=None, - bg_classifier_ratio=0.0, - loss_z1_factor=0.1, - weighting_mog=1.0, - warmup_model=False, # if true, trains a model without cell-type classification first. - ): - """Compute the loss.""" - px = generative_ouputs["px"] - qz1 = inference_outputs["qz"] - z1 = inference_outputs["z"] - x = tensors[REGISTRY_KEYS.X_KEY] - y = tensors[REGISTRY_KEYS.LABELS_KEY] - site_index = tensors[REGISTRY_KEYS.SITE_KEY] - - is_labelled = False if y is None else True - - # Enumerate choices of label - ys, z1s = broadcast_labels(z1, n_broadcast=self.n_labels_conditioning) - qz2, z2 = self.encoder_z2_z1(z1s, ys) - pz1_m, pz1_v = self.decoder_z1_z2(z2, ys) - reconst_loss = -px.log_prob(x).sum(-1) - - # KL Divergence - mean = torch.zeros_like(qz2.loc) - scale = torch.ones_like(qz2.scale) - - kl_divergence_z2 = kl(qz2, Normal(mean, scale)).sum(dim=1) - loss_z1_unweight = -Normal(pz1_m, torch.sqrt(pz1_v)).log_prob(z1s).sum(dim=-1) - loss_z1_weight = qz1.log_prob(z1).sum(dim=-1) - - probs, logits = self.classification(tensors, precomputed_z=z1) - probs_conditioning, logits_conditioning = ( - probs[self.conditioning_class], - logits[self.conditioning_class], - ) - - if z1.ndim == 2: - loss_z1_unweight_ = loss_z1_unweight.view(self.n_labels, -1).t() - kl_divergence_z2_ = kl_divergence_z2.view(self.n_labels, -1).t() - else: - loss_z1_unweight_ = torch.transpose( - loss_z1_unweight.view(z1.shape[0], self.n_labels, -1), -1, -2 - ) - kl_divergence_z2_ = torch.transpose( - kl_divergence_z2.view(z1.shape[0], self.n_labels, -1), -1, -2 - ) - reconst_loss += loss_z1_weight + (loss_z1_unweight_ * probs[-1]).sum(dim=-1) - kl_divergence = (kl_divergence_z2_ * probs[-1]).sum(dim=-1) - - if not warmup_model: - reconst_loss += ( - ( - loss_z1_weight - + ((loss_z1_unweight).view(self.n_labels_conditioning, -1).t() - * probs_conditioning).sum(dim=1) - ) - * kl_weight - * loss_z1_factor - ) - - if self.prior_z1 == "mog": - cats = Categorical(logits=self.prior_logits) - normal_dists = Independent( - Normal(self.prior_means, torch.exp(self.prior_log_scales) + 1e-4), - 1, - ) - prior = MixtureSameFamily(cats, normal_dists) - u = qz1.rsample(sample_shape=(30,)) - # (sample, n_obs, n_latent) -> (sample, n_obs,) - kl_divergence += -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) - elif self.prior_z1 == "mog_celltype": - if warmup_model: - # Assigns zero meaning equal weight to all unlabeled cells. Otherwise biases to sample from respective MoG. - logits_input = torch.nn.functional.one_hot( - y[:, self.mog_class].ravel().long(), self.num_classes[self.mog_class] + 1 - ).float()[:, :-1] - cats = Categorical(logits=10 * logits_input) - else: - cats = Categorical(logits=logits_conditioning) - normal_dists = torch.distributions.Independent( - Normal( - self.prior_z1_means.expand(x.shape[0], -1, -1), - torch.exp(self.prior_z1_scales).expand(x.shape[0], -1, -1) + 1e-2, - ), - reinterpreted_batch_ndims=1, - ) - - prior = MixtureSameFamily(cats, normal_dists) - u = qz1.rsample(sample_shape=(30,)) - # (sample, n_obs, n_latent) -> (sample, n_obs,) - kl_z = -(prior.log_prob(u) - qz1.log_prob(u).sum(-1)).mean(0) - kl_divergence += weighting_mog * kl_z - else: - prior = Normal(torch.zeros_like(qz1.loc), torch.ones_like(qz1.loc)) - kl_z = 0 - - probs_prior, _ = self.classification(tensors, precomputed_z=prior.sample()) - kl_divergence_cat = 0 - - if not warmup_model: - for i in range(0, self.total_level): - y_prior_ = ( - torch.stack( - [self.y_prior_fine[idx] for idx in site_index.ravel().long()], dim=0 - ) - if i == self.total_level - 1 - else getattr(self, f"y_prior_{i}") - ) - if self.update_yprior and i > 0: - if i == self.total_level - 1: - # Shape batch, coarse; batch, coarse, finer -> batch, finer - y_prior_ = torch.einsum("bc,bcf->bf", probs[i - 1], y_prior_) - else: - y_prior_ = torch.einsum("bc,cf->bf", probs[i - 1], y_prior_) - - kl_divergence_cat += kl( - Categorical(probs=probs[i]), - Categorical(probs=y_prior_), - ) - for i in range(0, self.total_level): - if i == self.total_level - 1: - y_prior_ = torch.stack( - [self.y_prior_fine[i] for i in site_index.ravel().long()], dim=0 - ) - else: - y_prior_ = self.__getattr__("y_prior_" + str(i)) - if self.update_yprior and i > 0: - if i == self.total_level - 1: - # Shape batch, coarse; batch, coarse, finer -> batch, finer - y_prior_ = torch.einsum("bc,bcf->bf", probs_prior[i - 1], y_prior_) - else: - y_prior_ = torch.einsum("bc,cf->bf", probs_prior[i - 1], y_prior_) - - kl_divergence_cat += bg_classifier_ratio * kl( - Categorical(probs=probs_prior[i]), - Categorical(probs=y_prior_), - ) - - kl_divergence += kl_divergence_cat - - loss = torch.mean(reconst_loss + kl_divergence * kl_weight) - - if labelled_tensors is not None: - # We filter cells with unlabeled_category in loss. - ce_loss, fine_true_labels, logits_fine = self.classification( - tensors, return_classifier_loss=True - ) - if not warmup_model: - loss += ce_loss * classification_ratio - return LossOutput( - loss=loss, - reconstruction_loss=reconst_loss, - kl_local=kl_divergence, - classification_loss=ce_loss, - true_labels=fine_true_labels, - logits=logits_fine, - ) - return LossOutput(loss=loss, reconstruction_loss=reconst_loss, kl_local=kl_divergence) diff --git a/src/scvi/external/muanvi/_utils.py b/src/scvi/external/muanvi/_utils.py deleted file mode 100644 index a93b9f12a3..0000000000 --- a/src/scvi/external/muanvi/_utils.py +++ /dev/null @@ -1,182 +0,0 @@ -import warnings -from collections.abc import Iterable as IterableClass -from collections.abc import Sequence - -import numpy as np -from anndata import AnnData -from pandas.api.types import CategoricalDtype - -from scvi import REGISTRY_KEYS, settings -from scvi.data import AnnDataManager -from scvi.data._utils import _make_column_categorical -from scvi.data.fields import CategoricalJointObsField - - -# Class creating a new Obsm field for partially annotated layers of labels -class LabelsWithUnlabeledJointObsField(CategoricalJointObsField): - """ - An AnnDataField for a collection of partially observed layers of labels .obs fields in the AnnData data structure. - - Creates an .obsm field compiling the given .obs fields. The model will reference the compiled - data as a whole. - - Parameters - ---------- - registry_key - Key to register field under in data registry. - attr_keys - Sequence of keys to combine to form the obsm or varm field. - unlabeled_category - A single category to represent unlabeled cells in the data. - """ - - MAPPINGS_KEY = "mappings" - FIELD_KEYS_KEY = "field_keys" - N_CATS_PER_KEY = "n_cats_per_key" - UNLABELED_CATEGORY = "unlabeled_category" - - def __init__( - self, - registry_key: str, - attr_keys: list[str] | None, - unlabeled_category: str | None, - ) -> None: - super().__init__(registry_key, attr_keys) - self.count_stat_key = f"n_{self.registry_key}" - self.unlabeled_category = unlabeled_category - - def _default_mappings_dict(self) -> dict: - return { - self.MAPPINGS_KEY: dict(), - self.FIELD_KEYS_KEY: [], - self.N_CATS_PER_KEY: [], - self.UNLABELED_CATEGORY: [], - } - - def _make_obsm_categorical( - self, adata: AnnData, category_dict: dict[str, list[str]] | None = None - ) -> dict: - if self.attr_keys != getattr(adata, self.attr_name)[self.attr_key].columns.tolist(): - raise ValueError( - f"Original .{self.source_attr_name} keys do not match the columns in the ", - f"generated .{self.attr_name} field.", - ) - - categories = {} - df = getattr(adata, self.attr_name)[self.attr_key] - for level, key in enumerate(self.attr_keys): - categorical_dtype = ( - CategoricalDtype(categories=category_dict[key]) - if category_dict is not None - else None - ) - if categorical_dtype is None: - categorical_obs = df[key].astype("category") - else: - categorical_obs = df[key].astype(categorical_dtype) - - mapping = categorical_obs.cat.categories.to_numpy(copy=True) - mapping = self._remap_unlabeled_to_final_category(mapping, level) - cat_dtype = CategoricalDtype(categories=mapping, ordered=True) - mapping = _make_column_categorical( - df, - key, - key, - categorical_dtype=cat_dtype, - warning=False - ) - categories[key] = mapping - - store_cats = categories if category_dict is None else category_dict - - mappings_dict = self._default_mappings_dict() - mappings_dict[self.MAPPINGS_KEY] = store_cats - mappings_dict[self.FIELD_KEYS_KEY] = self.attr_keys - mappings_dict[self.UNLABELED_CATEGORY] = self.unlabeled_category - for k in self.attr_keys: - mappings_dict[self.N_CATS_PER_KEY].append(len(store_cats[k]) - 1) - return mappings_dict - - def _remap_unlabeled_to_final_category(self, mapping: np.ndarray, level: int) -> np.ndarray: - # Make unlabeled category the last element - unlabeled_category = self.unlabeled_category - - # Check if the unlabeled category is in the mapping - if unlabeled_category in mapping: - # Find the index of the unlabeled category - unlabeled_idx = np.where(mapping == unlabeled_category)[0][0] - # Swap the unlabeled category with the last element - mapping[unlabeled_idx], mapping[-1] = mapping[-1], mapping[unlabeled_idx] - else: - # Append the unlabeled category if it's not in the mapping - mapping = np.append(mapping, unlabeled_category) - - return mapping - - def register_field(self, adata: AnnData) -> dict: - super().register_field(adata) - self._combine_fields(adata) - state_registry = self._make_obsm_categorical(adata) - return state_registry - - def transfer_field( - self, - state_registry: dict, - adata_target: AnnData, - extend_categories: bool = False, - allow_missing_labels: bool = False, - **kwargs, - ) -> dict: - """Transfer the field.""" - for level, key in enumerate(self.attr_keys): - if ( - allow_missing_labels - and key is not None - and key not in list(adata_target.obs.columns) - ): - # Fill in original .obs attribute with unlabeled_category values. - warnings.warn( - f"Missing labels key {key}. Filling in with " - f"unlabeled category {self.unlabeled_category}.", - UserWarning, - stacklevel=settings.warnings_stacklevel, - ) - adata_target.obs[key] = self.unlabeled_category[level] - - kwargs.pop("extend_categories", None) - transfer_state_registry = super().transfer_field( - state_registry, adata_target, extend_categories=extend_categories, **kwargs - ) - categories = {} - mapping = transfer_state_registry[self.MAPPINGS_KEY] - for level, key in enumerate(self.attr_keys): - mapping_ = self._remap_unlabeled_to_final_category(mapping[key], level) - categories[key] = mapping_ - store_cats = categories - mappings_dict = self._default_mappings_dict() - mappings_dict[self.MAPPINGS_KEY] = store_cats - mappings_dict[self.FIELD_KEYS_KEY] = self.attr_keys - mappings_dict[self.UNLABELED_CATEGORY] = self.unlabeled_category - for k in self.attr_keys: - mappings_dict[self.N_CATS_PER_KEY].append(len(store_cats[k]) - 1) - return mappings_dict - - -def _get_site_code_from_category(adata_manager: AnnDataManager, category: Sequence[int | str]): - if not isinstance(category, IterableClass) or isinstance(category, str): - category = [category] - - site_mappings = adata_manager.get_state_registry(REGISTRY_KEYS.SITE_KEY).categorical_mapping - site_code = [] - for cat in category: - if cat is None: - site_code.append(None) - continue - elif isinstance(cat, int) and cat < len(site_mappings): - site_code.append(site_mappings[cat]) - elif cat not in site_mappings: - raise ValueError(f'"{cat}" not a valid site category.') - else: - site_loc = np.where(site_mappings == cat)[0][0] - site_code.append(site_loc) - return site_code, site_mappings From a38562bfbb70ef0d9a266a7a5fb0cae45eaa3bc6 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 22 May 2026 09:17:13 +0000 Subject: [PATCH 15/24] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/scvi/external/__init__.py | 1 + src/scvi/external/scvix/__init__.py | 2 +- src/scvi/external/scvix/_model.py | 10 +-- src/scvi/external/scvix/_module.py | 89 +++++++++++++----------- src/scvi/model/base/_archesmixin.py | 2 +- src/scvi/model/base/_da_testing.py | 24 ++++--- src/scvi/model/base/_training_mixin.py | 1 - src/scvi/module/_vae.py | 2 +- src/scvi/module/base/_embedding_mixin.py | 22 +++--- src/scvi/module/base/_priors.py | 2 +- src/scvi/nn/_base_components.py | 53 ++++++++------ src/scvi/train/_trainingplans.py | 35 ++++++---- tests/external/muanvi/test_muanvi.py | 11 ++- tests/external/scvix/test_scvix.py | 12 +++- 14 files changed, 150 insertions(+), 116 deletions(-) diff --git a/src/scvi/external/__init__.py b/src/scvi/external/__init__.py index dfa15451d9..f6dcc18b17 100644 --- a/src/scvi/external/__init__.py +++ b/src/scvi/external/__init__.py @@ -2,6 +2,7 @@ from scvi import settings from scvi.utils import error_on_missing_dependencies + from .cellassign import CellAssign from .contrastivevi import ContrastiveVI from .cytovi import CYTOVI diff --git a/src/scvi/external/scvix/__init__.py b/src/scvi/external/scvix/__init__.py index aea7c60506..aa49a88060 100644 --- a/src/scvi/external/scvix/__init__.py +++ b/src/scvi/external/scvix/__init__.py @@ -1,4 +1,4 @@ from ._model import SCVIX from ._module import VAEX -__all__ = ["SCVIX", "VAEX"] \ No newline at end of file +__all__ = ["SCVIX", "VAEX"] diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index a4dd7106af..c1247aa1f0 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -282,7 +282,7 @@ def train( else: adversarial_classifier = False n_epochs_kl_warmup = ( - n_epochs_kl_warmup if n_epochs_kl_warmup is not None else max_epochs//2 + n_epochs_kl_warmup if n_epochs_kl_warmup is not None else max_epochs // 2 ) if reduce_lr_on_plateau: check_val_every_n_epoch = 1 @@ -372,11 +372,13 @@ def setup_anndata( if labels_key is not None: anndata_fields.append( LabelsWithUnlabeledObsField( - REGISTRY_KEYS.LABELS_KEY, labels_key, unlabeled_category)) + REGISTRY_KEYS.LABELS_KEY, labels_key, unlabeled_category + ) + ) if adversarial_group_key is not None: anndata_fields.append( - CategoricalObsField( - REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, adversarial_group_key)) + CategoricalObsField(REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, adversarial_group_key) + ) # register new fields if the adata is minified adata_minify_type = _get_adata_minify_type(adata) if adata_minify_type is not None: diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index a70230820e..a314a09c21 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -14,8 +14,8 @@ from scvi.module.base import ( BaseMinifiedModeModuleClass, EmbeddingModuleMixin, - LossOutput, GaussianPrior, + LossOutput, MogPrior, VampPrior, auto_move_data, @@ -174,13 +174,12 @@ def __init__( self.use_observed_lib_size = True self.n_hidden = n_hidden - if self.dispersion == "gene": - self.px_r = torch.nn.Parameter(3.*torch.ones(n_input)) + self.px_r = torch.nn.Parameter(3.0 * torch.ones(n_input)) elif self.dispersion == "gene-batch": - self.px_r = torch.nn.Parameter(3.*torch.ones(n_input, n_batch)) + self.px_r = torch.nn.Parameter(3.0 * torch.ones(n_input, n_batch)) elif self.dispersion == "gene-assay": - self.px_r = torch.nn.Parameter(3.*torch.ones(n_input, n_assay)) + self.px_r = torch.nn.Parameter(3.0 * torch.ones(n_input, n_assay)) elif self.dispersion == "gene-cell": pass else: @@ -284,10 +283,7 @@ def __init__( assert pseudoinput_data is not None, ( "Pseudoinput data must be specified if using VampPrior" ) - pseudoinput_data = self._get_inference_input( - pseudoinput_data, - full_forward_pass=True - ) + pseudoinput_data = self._get_inference_input(pseudoinput_data, full_forward_pass=True) cat_list = [n_batch] + n_cats_per_cov_ + encode_assay_list self.prior = VampPrior( n_components=n_prior_components, @@ -296,7 +292,7 @@ def __init__( pseudoinputs=pseudoinput_data, n_cat_list=cat_list, trainable_priors=True, - additional_categorical_covariates=["assay_index"] + additional_categorical_covariates=["assay_index"], ) elif prior == "mog": self.prior = MogPrior( @@ -304,14 +300,9 @@ def __init__( n_latent=n_latent, ) elif prior == "mog_celltype": - self.prior = MogPrior( - n_components=n_labels, - n_latent=n_latent, - celltype_bias=True - ) + self.prior = MogPrior(n_components=n_labels, n_latent=n_latent, celltype_bias=True) else: - raise ValueError( - "`prior` must be one of 'gaussian', 'vamp', 'mog', 'mog_celltype'.") + raise ValueError("`prior` must be one of 'gaussian', 'vamp', 'mog', 'mog_celltype'.") def _get_inference_input( self, @@ -336,7 +327,9 @@ def _get_inference_input( MODULE_KEYS.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), 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.ADVERSARIAL_GROUP_KEY: tensors.get(REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, None), + MODULE_KEYS.ADVERSARIAL_GROUP_KEY: tensors.get( + REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, None + ), } else: return { @@ -376,7 +369,7 @@ def _regular_inference( if self.use_observed_lib_size: library = torch.log(x.sum(1)).unsqueeze(1) if self.log_variational: - x_ = x_/x_.mean(1).unsqueeze(1) + x_ = x_ / x_.mean(1).unsqueeze(1) x_ = torch.log1p(x_) if cat_covs is not None and self.encode_covariates: @@ -402,9 +395,7 @@ def _regular_inference( if n_samples > 1: untran_z = qz.sample((n_samples,)) z = self.z_encoder.z_transformation(untran_z) - library = library.unsqueeze(0).expand( - (n_samples, library.size(0), library.size(1)) - ) + library = library.unsqueeze(0).expand((n_samples, library.size(0), library.size(1))) return { MODULE_KEYS.Z_KEY: z, @@ -447,7 +438,7 @@ def generative( assay_index: torch.Tensor | None = None, cont_covs: torch.Tensor | None = None, cat_covs: torch.Tensor | None = None, - size_factor: torch.Tensor | None = None, # Consistency + size_factor: torch.Tensor | None = None, # Consistency y: torch.Tensor | None = None, transform_batch: torch.Tensor | None = None, transform_assay: torch.Tensor | None = None, @@ -544,7 +535,7 @@ def loss( weight_assay_loss: float = 0.0, weight_global: float = 1.0, weight_kl_sample: float = 1.0, - classification_ratio: float = 500., + classification_ratio: float = 500.0, ) -> LossOutput: """Compute the loss.""" from torch.distributions import kl_divergence @@ -560,9 +551,11 @@ def loss( if self.get_embedding_variational(REGISTRY_KEYS.BATCH_KEY, default_value=False): pz_sample = self.compute_embedding( - REGISTRY_KEYS.BATCH_KEY, tensors[REGISTRY_KEYS.BATCH_KEY], return_dist=True) + REGISTRY_KEYS.BATCH_KEY, tensors[REGISTRY_KEYS.BATCH_KEY], return_dist=True + ) qz_sample = distributions.Normal( - torch.zeros_like(pz_sample.loc), torch.ones_like(pz_sample.scale)) + torch.zeros_like(pz_sample.loc), torch.ones_like(pz_sample.scale) + ) kl_divergence_sample = kl_divergence(qz_sample, pz_sample).sum(dim=1) else: kl_divergence_sample = torch.zeros_like(kl_divergence_z) @@ -570,8 +563,7 @@ def loss( assay_index = tensors.get(REGISTRY_KEYS.ASSAY_KEY, None) if weight_assay_loss > 0.0 and assay_index is not None: assay_loss = self._compute_assay_penalty( - inference_outputs[MODULE_KEYS.QZ_KEY].loc, - assay_index + inference_outputs[MODULE_KEYS.QZ_KEY].loc, assay_index ) else: assay_loss = 0.0 @@ -582,30 +574,44 @@ def loss( if self.gene_likelihood == "zinb": zi = generative_outputs[MODULE_KEYS.PX_KEY].zi_logits - kl_global -= distributions.Exponential( - 10.*torch.ones_like(zi)).log_prob(torch.exp(zi)).sum(-1) + kl_global -= ( + distributions.Exponential(10.0 * torch.ones_like(zi)) + .log_prob(torch.exp(zi)) + .sum(-1) + ) if self.gene_likelihood == "zinb" or self.gene_likelihood == "nb": theta = generative_outputs[MODULE_KEYS.PX_KEY].theta - kl_global -= distributions.Exponential( - torch.ones_like(theta)).log_prob(1/theta).sum(-1) + kl_global -= ( + distributions.Exponential(torch.ones_like(theta)).log_prob(1 / theta).sum(-1) + ) weighted_kl_local = kl_weight * (kl_divergence_z + weight_kl_sample * kl_divergence_sample) loss = torch.mean( - reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss + - weight_global * kl_global) + reconst_loss + + weighted_kl_local + + weight_assay_loss * assay_loss + + weight_global * kl_global + ) if self.n_labels > 1: logits = self.classifier(inference_outputs[MODULE_KEYS.Z_KEY]) - classification_loss_ = torch.nn.functional.cross_entropy(logits, y.ravel(), reduction="none") - mask = (y != self.n_labels) + classification_loss_ = torch.nn.functional.cross_entropy( + logits, y.ravel(), reduction="none" + ) + mask = y != self.n_labels classification_loss = classification_ratio * torch.masked_select( - classification_loss_, mask).mean(0) + classification_loss_, mask + ).mean(0) loss += torch.mean(classification_loss) return LossOutput( - loss=loss, reconstruction_loss=reconst_loss, kl_local=kl_divergence_z, - classification_loss=classification_loss, logits=logits, true_labels=y + loss=loss, + reconstruction_loss=reconst_loss, + kl_local=kl_divergence_z, + classification_loss=classification_loss, + logits=logits, + true_labels=y, ) return LossOutput( @@ -667,11 +673,10 @@ def sample( return samples.cpu() - def _compute_assay_penalty( - self, params, assay): + def _compute_assay_penalty(self, params, assay): assay = assay.squeeze(-1).long() unique = torch.unique(assay) - pair_penalty = torch.tensor(0., device=assay.device) + pair_penalty = torch.tensor(0.0, device=assay.device) if len(unique) > 1: for i in unique: pp = self.mmd(params, mask=(assay == i)) diff --git a/src/scvi/model/base/_archesmixin.py b/src/scvi/model/base/_archesmixin.py index f9654b4a65..25e1b6e1fb 100644 --- a/src/scvi/model/base/_archesmixin.py +++ b/src/scvi/model/base/_archesmixin.py @@ -536,4 +536,4 @@ def _pad_and_sort_query_anndata( if adata_out is not adata: adata._init_as_actual(adata_out) else: - return adata_out \ No newline at end of file + return adata_out diff --git a/src/scvi/model/base/_da_testing.py b/src/scvi/model/base/_da_testing.py index e50e701153..cdf183b15b 100644 --- a/src/scvi/model/base/_da_testing.py +++ b/src/scvi/model/base/_da_testing.py @@ -14,7 +14,7 @@ def get_aggregated_posterior( sample: str | int | None = None, indices: Sequence[int] | None = None, batch_size: int | None = None, - dof: float | None = 3., + dof: float | None = 3.0, ) -> dist.Distribution: """Compute the aggregated posterior over the ``u`` latent representations. @@ -46,17 +46,21 @@ def get_aggregated_posterior( ) dataloader = self._make_data_loader(adata=adata, indices=indices, batch_size=batch_size) - qu_loc, qu_scale = self.get_latent_representation(batch_size=batch_size, return_dist=True, dataloader=dataloader, give_mean=True) + qu_loc, qu_scale = self.get_latent_representation( + batch_size=batch_size, return_dist=True, dataloader=dataloader, give_mean=True + ) - qu_loc = torch.tensor(qu_loc, device='cuda').T - qu_scale = torch.tensor(qu_scale, device='cuda').T + qu_loc = torch.tensor(qu_loc, device="cuda").T + qu_scale = torch.tensor(qu_scale, device="cuda").T if dof is None: components = dist.Normal(qu_loc, qu_scale) else: components = dist.StudentT(dof, qu_loc, qu_scale) return dist.MixtureSameFamily( - dist.Categorical(logits=torch.ones(qu_loc.shape[1], device='cuda')), components) + dist.Categorical(logits=torch.ones(qu_loc.shape[1], device="cuda")), components + ) + def differential_abundance( self, @@ -92,9 +96,7 @@ def differential_abundance( """ adata = self._validate_anndata(adata) - us = self.get_latent_representation( - batch_size=batch_size, return_dist=False, give_mean=True - ) + us = self.get_latent_representation(batch_size=batch_size, return_dist=False, give_mean=True) unique_samples = adata.obs[sample_key].unique() dataloader = torch.utils.data.DataLoader(us, batch_size=batch_size) @@ -107,10 +109,12 @@ def differential_abundance( ap = get_aggregated_posterior(self, adata=adata, indices=indices, dof=dof) log_probs_ = [] for u_rep in dataloader: - u_rep = u_rep.to('cuda') + u_rep = u_rep.to("cuda") log_probs_.append(ap.log_prob(u_rep).sum(-1, keepdims=True)) log_probs.append(torch.cat(log_probs_, axis=0).cpu().numpy()) log_probs = np.concatenate(log_probs, 1) - log_probs_df = pd.DataFrame(data=log_probs, index=adata.obs_names.to_numpy(), columns=unique_samples) + log_probs_df = pd.DataFrame( + data=log_probs, index=adata.obs_names.to_numpy(), columns=unique_samples + ) return log_probs_df diff --git a/src/scvi/model/base/_training_mixin.py b/src/scvi/model/base/_training_mixin.py index 52d09cc054..cc7d935891 100644 --- a/src/scvi/model/base/_training_mixin.py +++ b/src/scvi/model/base/_training_mixin.py @@ -14,7 +14,6 @@ from scvi.dataloaders import DataSplitter, SemiSupervisedDataSplitter from scvi.model._utils import get_max_epochs_heuristic, use_distributed_sampler from scvi.train import ( - AdversarialTrainingPlan, SemiSupervisedAdversarialTrainingPlan, SemiSupervisedTrainingPlan, TrainingPlan, diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index fae77d8242..8f74ab4f09 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -370,7 +370,7 @@ def _regular_inference( if self.use_observed_lib_size: library = torch.log(x.sum(1)).unsqueeze(1) if self.log_variational: - x_ = x_/x_.mean(1).unsqueeze(1) + x_ = x_ / x_.mean(1).unsqueeze(1) x_ = torch.log1p(x_) if cont_covs is not None and self.encode_covariates: diff --git a/src/scvi/module/base/_embedding_mixin.py b/src/scvi/module/base/_embedding_mixin.py index 7651358023..02dd4ebef9 100644 --- a/src/scvi/module/base/_embedding_mixin.py +++ b/src/scvi/module/base/_embedding_mixin.py @@ -43,7 +43,10 @@ def remove_embedding(self, key: str) -> None: raise KeyError(f"Embedding {key} not found.") del self.embeddings_dict[key] - def get_embedding(self, key: str,) -> Embedding: + def get_embedding( + self, + key: str, + ) -> Embedding: """Get an embedding from the module.""" if key not in self.embeddings_dict: raise KeyError(f"Embedding {key} not found.") @@ -85,12 +88,12 @@ def init_embedding( @auto_move_data def compute_embedding( - self, - key: str, - indices: torch.Tensor, - return_mean: bool = False, - return_dist: bool = False, - ) -> torch.Tensor: + self, + key: str, + indices: torch.Tensor, + return_mean: bool = False, + return_dist: bool = False, + ) -> torch.Tensor: """Forward pass for an embedding.""" indices = indices.flatten() if indices.ndim > 1 else indices embedding = self.get_embedding(key)(indices) @@ -98,10 +101,7 @@ def compute_embedding( embedding_dim = self.get_embedding_dim(key) if return_mean: return embedding[:, :embedding_dim] - dist = Normal( - embedding[:, :embedding_dim], - torch.exp(embedding[:, embedding_dim:]) - ) + dist = Normal(embedding[:, :embedding_dim], torch.exp(embedding[:, embedding_dim:])) if return_dist: return dist else: diff --git a/src/scvi/module/base/_priors.py b/src/scvi/module/base/_priors.py index 5fd0dcab8d..6778c908a5 100644 --- a/src/scvi/module/base/_priors.py +++ b/src/scvi/module/base/_priors.py @@ -96,7 +96,7 @@ def __init__( ): super().__init__() self.prior_means = torch.nn.Parameter(0.1 * torch.randn([n_components, n_latent])) - self.prior_log_scales = torch.nn.Parameter(torch.zeros([n_components, n_latent]) - 1.) + self.prior_log_scales = torch.nn.Parameter(torch.zeros([n_components, n_latent]) - 1.0) self.prior_logits = torch.nn.Parameter(torch.zeros([n_components])) self.celltype_bias = celltype_bias if celltype_bias: diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index dec01f161e..faf5ebb5a0 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -20,8 +20,10 @@ def __init__(self, num_features, num_classes, momentum, eps): self.num_features = num_features self.bn = nn.BatchNorm1d(self.num_features, momentum=momentum, eps=eps, affine=False) self.embed = nn.Embedding(num_classes, self.num_features * 2) - self.embed.weight.data[:, :self.num_features].normal_(1, 0.02) # Initialise scale at N(1, 0.02) - self.embed.weight.data[:, self.num_features:].zero_() # Initialise bias at 0 + self.embed.weight.data[:, : self.num_features].normal_( + 1, 0.02 + ) # Initialise scale at N(1, 0.02) + self.embed.weight.data[:, self.num_features :].zero_() # Initialise bias at 0 def forward(self, x, y): out = self.bn(x) @@ -30,14 +32,17 @@ def forward(self, x, y): return out + class ConditionalLayerNorm(nn.Module): def __init__(self, num_features, num_classes): super().__init__() self.num_features = num_features self.ln = nn.LayerNorm(self.num_features, elementwise_affine=False) self.embed = nn.Embedding(num_classes, self.num_features * 2) - self.embed.weight.data[:, :self.num_features].normal_(1, 0.02) # Initialise scale at N(1, 0.02) - self.embed.weight.data[:, self.num_features:].zero_() # Initialise bias at 0 + self.embed.weight.data[:, : self.num_features].normal_( + 1, 0.02 + ) # Initialise scale at N(1, 0.02) + self.embed.weight.data[:, self.num_features :].zero_() # Initialise bias at 0 def forward(self, x, y): out = self.ln(x) @@ -46,6 +51,7 @@ def forward(self, x, y): return out + class FCLayers(nn.Module): """A helper class to build fully-connected layers for a neural network. @@ -100,7 +106,7 @@ def __init__( inject_covariates: bool = True, activation_fn: nn.Module = nn.ReLU, conditional_norm: bool = False, - conditional_category: int = 0 + conditional_category: int = 0, ): super().__init__() self.inject_covariates = inject_covariates @@ -114,7 +120,7 @@ def __init__( self.n_continuous = n_continuous self.cond_cat = conditional_category - if conditional_norm and self.n_cat_list[self.cond_cat]==0: + if conditional_norm and self.n_cat_list[self.cond_cat] == 0: raise ValueError( "Conditional normalization is not applicable for a categorical variable with only one category." ) @@ -134,14 +140,17 @@ def __init__( ), # non-default params come from defaults in Tensorflow implementation ConditionalBatchNorm2d( - n_out, self.n_cat_list[self.cond_cat], momentum=0.01, eps=0.001) + n_out, self.n_cat_list[self.cond_cat], momentum=0.01, eps=0.001 + ) if conditional_norm and use_batch_norm - else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) if use_batch_norm + else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) + if use_batch_norm else None, # non-default params come from defaults in Tensorflow implementation ConditionalLayerNorm(n_out, self.n_cat_list[self.cond_cat]) if conditional_norm and use_layer_norm - else nn.LayerNorm(n_out, elementwise_affine=False) if use_layer_norm + else nn.LayerNorm(n_out, elementwise_affine=False) + if use_layer_norm else None, activation_fn() if use_activation else None, nn.Dropout(p=dropout_rate) if dropout_rate > 0 else None, @@ -209,7 +218,7 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No if len(self.n_cat_list) > len(cat_list): raise ValueError("nb. categorical args provided doesn't match init. params.") - if self.n_continuous>0 and cont_input.shape[-1] != self.n_continuous: + if self.n_continuous > 0 and cont_input.shape[-1] != self.n_continuous: raise ValueError("continuous dims provided doesn't match init. params.") for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): if n_cat and cat is None: @@ -226,12 +235,15 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No for layer in layers: if layer is not None: if isinstance(layer, ConditionalBatchNorm2d) or isinstance( - layer, ConditionalLayerNorm): + layer, ConditionalLayerNorm + ): if x.dim() == 3: x = torch.cat( - [(layer(x=slice_x, y=cat_list[self.cond_cat])).unsqueeze(0) - for slice_x in x], - dim=0 + [ + (layer(x=slice_x, y=cat_list[self.cond_cat])).unsqueeze(0) + for slice_x in x + ], + dim=0, ) else: x = layer(x=x, y=cat_list[self.cond_cat]) @@ -342,11 +354,11 @@ def __init__( self.var_activation = torch.exp if var_activation is None else var_activation def forward( - self, - x: torch.Tensor, - *cat_list: int, - cont: torch.Tensor | None = None, - ): + self, + x: torch.Tensor, + *cat_list: int, + cont: torch.Tensor | None = None, + ): r"""The forward computation for a single sample. #. Encodes the data into latent space using the encoder network @@ -512,8 +524,7 @@ def forward( px = self.px_decoder(z, *cat_list, cont=cont) if output_condition is not None and self.n_conditions_output: one_hot_cat = nn.functional.one_hot( - output_condition.squeeze(-1), - self.n_conditions_output + output_condition.squeeze(-1), self.n_conditions_output ) else: one_hot_cat = torch.zeros(px.size(-2), self.n_conditions_output).to(px.device) diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 8fa45b2236..d47c2ea4d0 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -17,13 +17,12 @@ from scvi import REGISTRY_KEYS, settings from scvi.module import Classifier -from scvi.module._constants import MODULE_KEYS from scvi.module.base import ( BaseModuleClass, LossOutput, MogPrior, - VampPrior, PyroBaseModuleClass, + VampPrior, ) from scvi.train._constants import METRIC_KEYS from scvi.utils import is_package_installed @@ -655,7 +654,7 @@ def __init__( self.adversarial_classifier = False else: self.adversarial_classifier = Classifier( - n_input=self.module.n_latent+self.module.n_adversarial_group, + n_input=self.module.n_latent + self.module.n_adversarial_group, n_hidden=128, n_labels=self.n_output_classifier, n_layers=2, @@ -669,13 +668,15 @@ def __init__( self.scale_adversarial_loss = scale_adversarial_loss self.automatic_optimization = False - def loss_adversarial_classifier(self, z, adversarial_group, batch_index, predict_true_class=True): + def loss_adversarial_classifier( + self, z, adversarial_group, batch_index, predict_true_class=True + ): """Loss for adversarial classifier.""" n_classes = self.n_output_classifier adversarial_group_ = torch.nn.functional.one_hot( adversarial_group, num_classes=self.module.n_adversarial_group ).float() - if predict_true_class: # train classifier + if predict_true_class: # train classifier z = z.detach() z = torch.cat([z, adversarial_group_], dim=1) cls_logits = self.adversarial_classifier(z) @@ -686,9 +687,11 @@ def loss_adversarial_classifier(self, z, adversarial_group, batch_index, predict else: one_hot_batch = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes).float() cls_target = (1 - one_hot_batch) / (n_classes - 1) - loss = - ( - cls_target * torch.nn.functional.log_softmax(cls_logits, dim=1) - ).sum(dim=1).mean() + loss = ( + -(cls_target * torch.nn.functional.log_softmax(cls_logits, dim=1)) + .sum(dim=1) + .mean() + ) return loss @@ -739,19 +742,23 @@ def training_step(self, batch, batch_idx): # train adversarial classifier # this condition will not be met unless self.adversarial_classifier is not False if opt2 is not None: - loss = 0. + loss = 0.0 for i in range(self.adversarial_steps): qz = inference_outputs["qz"] z = qz.sample() - loss_ = kappa * self.loss_adversarial_classifier(z, adversarial_group, batch_tensor, True) - if isinstance(self.module.prior, MogPrior) or isinstance(self.module.prior, VampPrior): + loss_ = kappa * self.loss_adversarial_classifier( + z, adversarial_group, batch_tensor, True + ) + if isinstance(self.module.prior, MogPrior) or isinstance( + self.module.prior, VampPrior + ): qz_m, qz_v = qz.loc.detach(), qz.scale.detach() loss_ += self.module.prior.kl( qz=Normal(qz_m, qz_v), z=z, labels=batch.get(REGISTRY_KEYS.LABELS_KEY, torch.tensor(0)).long(), ).mean() - if i>1 and (loss - loss_)/loss < 1e-3: + if i > 1 and (loss - loss_) / loss < 1e-3: break loss = loss_ opt2.zero_grad() @@ -805,9 +812,7 @@ def configure_optimizers(self): if self.adversarial_classifier is not False: params2 = filter(lambda p: p.requires_grad, self.adversarial_classifier.parameters()) - optimizer2 = torch.optim.Adam( - params2, lr=3e-4, eps=1e-4, weight_decay=1e-9 - ) + optimizer2 = torch.optim.Adam(params2, lr=3e-4, eps=1e-4, weight_decay=1e-9) config2 = {"optimizer": optimizer2} # pytorch lightning requires this way to return diff --git a/tests/external/muanvi/test_muanvi.py b/tests/external/muanvi/test_muanvi.py index 5d64563866..1d2cba0e65 100644 --- a/tests/external/muanvi/test_muanvi.py +++ b/tests/external/muanvi/test_muanvi.py @@ -1,14 +1,13 @@ -import pytest -from mudata import MuData -import numpy as np -from anndata import AnnData import itertools -from scipy import sparse as sp_sparse +import numpy as np import pandas as pd +import pytest +from anndata import AnnData +from mudata import MuData +from scipy import sparse as sp_sparse from scvi.data import synthetic_iid -from scvi.external import MUANVI # helper function for testing purposes ; could be moved in the same file as generate_synthetic() diff --git a/tests/external/scvix/test_scvix.py b/tests/external/scvix/test_scvix.py index 89ec21a764..0d926ffce1 100644 --- a/tests/external/scvix/test_scvix.py +++ b/tests/external/scvix/test_scvix.py @@ -1,13 +1,14 @@ -import pytest import os + import numpy as np +import pytest import scvi from scvi.data import synthetic_iid from scvi.data._constants import ADATA_MINIFY_TYPE from scvi.data._utils import _is_minified -from scvi.model.base import BaseMinifiedModeModelClass from scvi.external import SCVIX +from scvi.model.base import BaseMinifiedModeModelClass def assert_approx_equal(a, b): @@ -30,6 +31,7 @@ def test_scvix(prior: str): model.get_reconstruction_error(indices=model.validation_indices) model.differential_expression(groupby="labels", group1="label_1") + @pytest.mark.parametrize("dispersion", ["gene", "gene-batch", "gene-assay", "gene-cell"]) def test_scvix_dispersion(dispersion: str): adata = synthetic_iid(batch_size=100) @@ -38,6 +40,7 @@ def test_scvix_dispersion(dispersion: str): model.train(max_epochs=1) model.get_normalized_expression() + def test_scvix_encode_covariates(): adata = synthetic_iid(batch_size=100) SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") @@ -45,6 +48,7 @@ def test_scvix_encode_covariates(): model.train(max_epochs=1) model.get_normalized_expression(n_samples=2) + def test_scvix_embedding(): adata = synthetic_iid(batch_size=100) SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") @@ -52,6 +56,7 @@ def test_scvix_embedding(): model.train(max_epochs=1) model.get_normalized_expression(n_samples=2) + def test_scvix_layernorm(): adata = synthetic_iid(batch_size=100) SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") @@ -62,6 +67,7 @@ def test_scvix_layernorm(): model.train(max_epochs=1) model.get_normalized_expression(n_samples=2) + def test_scvix_scarches_one_hot(save_path): # test transfer_anndata_setup + view adata1 = synthetic_iid() @@ -88,6 +94,7 @@ def test_scvix_scarches_one_hot(save_path): new_var_names = new_var_names_init + adata4.var_names[10:].to_list() adata4.var_names = new_var_names + def test_scvix_scarches_embedding(save_path): # test transfer_anndata_setup + view adata1 = synthetic_iid() @@ -114,6 +121,7 @@ def test_scvix_scarches_embedding(save_path): new_var_names = new_var_names_init + adata4.var_names[10:].to_list() adata4.var_names = new_var_names + def test_scvix_minified(): adata = synthetic_iid() SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") From 6934058ecfd7a583b62e7b3693f2b707394965de Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Fri, 22 May 2026 11:35:04 +0200 Subject: [PATCH 16/24] Remove changes from rebase. --- src/scvi/_constants.py | 1 - src/scvi/data/_utils.py | 3 +- src/scvi/data/fields/_scanvi.py | 1 - src/scvi/model/_scvi.py | 4 +- src/scvi/model/base/_da_testing.py | 120 ---------------------- src/scvi/module/_constants.py | 1 - src/scvi/module/_vae.py | 8 +- src/scvi/nn/_base_components.py | 4 +- src/scvi/nn/_embedding.py | 6 +- src/scvi/train/_trainingplans.py | 11 -- src/scvi/utils/_docstrings.py | 2 +- tests/external/muanvi/test_muanvi.py | 144 --------------------------- 12 files changed, 14 insertions(+), 291 deletions(-) delete mode 100644 src/scvi/model/base/_da_testing.py delete mode 100644 tests/external/muanvi/test_muanvi.py diff --git a/src/scvi/_constants.py b/src/scvi/_constants.py index 230a897037..28caa685f1 100644 --- a/src/scvi/_constants.py +++ b/src/scvi/_constants.py @@ -5,7 +5,6 @@ class _REGISTRY_KEYS_NT(NamedTuple): X_KEY: str = "X" ATAC_X_KEY: str = "atac" BATCH_KEY: str = "batch" - SITE_KEY: str = "site" ASSAY_KEY: str = "assay" ADVERSARIAL_GROUP_KEY: str = "adversarial_group" SAMPLE_KEY: str = "sample" diff --git a/src/scvi/data/_utils.py b/src/scvi/data/_utils.py index f1af020e91..1b5c5a5253 100644 --- a/src/scvi/data/_utils.py +++ b/src/scvi/data/_utils.py @@ -197,7 +197,6 @@ def _make_column_categorical( column_key: str, alternate_column_key: str, categorical_dtype: str | CategoricalDtype | None = None, - warning: bool = True, ): """Makes the data in column_key in DataFrame all categorical. @@ -222,7 +221,7 @@ def _make_column_categorical( df[alternate_column_key] = codes # make sure each category contains enough cells - if np.min(counts) < 3 and warning: + if np.min(counts) < 3: category = unique[np.argmin(counts)] warnings.warn( f"Category {category} in adata.obs['{alternate_column_key}'] has fewer than 3 cells. " diff --git a/src/scvi/data/fields/_scanvi.py b/src/scvi/data/fields/_scanvi.py index db3dcfb85b..6e1d443b5a 100644 --- a/src/scvi/data/fields/_scanvi.py +++ b/src/scvi/data/fields/_scanvi.py @@ -58,7 +58,6 @@ def _remap_unlabeled_to_final_category(self, adata: AnnData, mapping: np.ndarray self._original_attr_key, self.attr_key, categorical_dtype=cat_dtype, - warning=False, ) return { diff --git a/src/scvi/model/_scvi.py b/src/scvi/model/_scvi.py index d527dfd554..e596883e07 100644 --- a/src/scvi/model/_scvi.py +++ b/src/scvi/model/_scvi.py @@ -180,7 +180,9 @@ def __init__( n_cats_per_cov = None n_batch = self.summary_stats.n_batch - use_size_factor_key = REGISTRY_KEYS.SIZE_FACTOR_KEY in self.adata_manager.data_registry + use_size_factor_key = self.registry_["setup_args"][ + f"{REGISTRY_KEYS.SIZE_FACTOR_KEY}_key" + ] library_log_means, library_log_vars = None, None if ( not use_size_factor_key diff --git a/src/scvi/model/base/_da_testing.py b/src/scvi/model/base/_da_testing.py deleted file mode 100644 index cdf183b15b..0000000000 --- a/src/scvi/model/base/_da_testing.py +++ /dev/null @@ -1,120 +0,0 @@ -from collections.abc import Sequence - -import numpy as np -import pandas as pd -import torch -import torch.distributions as dist -from anndata import AnnData -from tqdm import tqdm - - -def get_aggregated_posterior( - self, - adata: AnnData | None = None, - sample: str | int | None = None, - indices: Sequence[int] | None = None, - batch_size: int | None = None, - dof: float | None = 3.0, -) -> dist.Distribution: - """Compute the aggregated posterior over the ``u`` latent representations. - - Parameters - ---------- - adata - AnnData object to use. Defaults to the AnnData object used to initialize the model. - sample - Name or index of the sample to filter on. If ``None``, uses all cells. - indices - Indices of cells to use. - batch_size - Batch size to use for computing the latent representation. - dof - Degrees of freedom for the Student's t-distribution components. If ``None``, components are Normal. - - Returns - ------- - A mixture distribution of the aggregated posterior. - """ - self._check_if_trained(warn=False) - adata = self._validate_anndata(adata) - - if indices is None: - indices = np.arange(self.adata.n_obs) - if sample is not None: - indices = np.intersect1d( - np.array(indices), np.where(adata.obs[self.sample_key] == sample)[0] - ) - - dataloader = self._make_data_loader(adata=adata, indices=indices, batch_size=batch_size) - qu_loc, qu_scale = self.get_latent_representation( - batch_size=batch_size, return_dist=True, dataloader=dataloader, give_mean=True - ) - - qu_loc = torch.tensor(qu_loc, device="cuda").T - qu_scale = torch.tensor(qu_scale, device="cuda").T - - if dof is None: - components = dist.Normal(qu_loc, qu_scale) - else: - components = dist.StudentT(dof, qu_loc, qu_scale) - return dist.MixtureSameFamily( - dist.Categorical(logits=torch.ones(qu_loc.shape[1], device="cuda")), components - ) - - -def differential_abundance( - self, - adata: AnnData | None = None, - sample_key: str | None = None, - batch_size: int = 128, - downsample_cells: int | None = None, - dof: float | None = None, -) -> pd.DataFrame: - """Compute the differential abundance between samples. - - Computes the logarithm of the ratio of the probabilities of each sample conditioned on the - estimated aggregate posterior distribution of each cell. - - Parameters - ---------- - adata - The data object to compute the differential abundance for. - sample_key - Key for the sample covariate. - batch_size - Minibatch size for computing the differential abundance. - downsample_cells - Number of cells to subset to before computing the differential abundance. - dof - Degrees of freedom for the Student's t-distribution components for aggregated posterior. If ``None``, components are Normal. - - Returns - ------- - DataFrame of shape (n_cells, n_samples) containing the log probabilities - for each cell across samples. The rows correspond to cell names from `adata.obs_names`, - and the columns correspond to unique sample identifiers. - """ - adata = self._validate_anndata(adata) - - us = self.get_latent_representation(batch_size=batch_size, return_dist=False, give_mean=True) - - unique_samples = adata.obs[sample_key].unique() - dataloader = torch.utils.data.DataLoader(us, batch_size=batch_size) - log_probs = [] - for sample_name in tqdm(unique_samples): - indices = np.where(adata.obs[sample_key] == sample_name)[0] - if downsample_cells is not None and downsample_cells < indices.shape[0]: - indices = np.random.choice(indices, downsample_cells, replace=False) - - ap = get_aggregated_posterior(self, adata=adata, indices=indices, dof=dof) - log_probs_ = [] - for u_rep in dataloader: - u_rep = u_rep.to("cuda") - log_probs_.append(ap.log_prob(u_rep).sum(-1, keepdims=True)) - log_probs.append(torch.cat(log_probs_, axis=0).cpu().numpy()) - - log_probs = np.concatenate(log_probs, 1) - log_probs_df = pd.DataFrame( - data=log_probs, index=adata.obs_names.to_numpy(), columns=unique_samples - ) - return log_probs_df diff --git a/src/scvi/module/_constants.py b/src/scvi/module/_constants.py index 6626022bea..885dc4fe68 100644 --- a/src/scvi/module/_constants.py +++ b/src/scvi/module/_constants.py @@ -13,7 +13,6 @@ class _MODULE_KEYS(NamedTuple): BATCH_INDEX_KEY: str = "batch_index" ASSAY_INDEX_KEY: str = "assay_index" ADVERSARIAL_GROUP_KEY: str = "adversarial_group" - SITE_INDEX_KEY: str = "site_index" Y_KEY: str = "y" CONT_COVS_KEY: str = "cont_covs" CAT_COVS_KEY: str = "cat_covs" diff --git a/src/scvi/module/_vae.py b/src/scvi/module/_vae.py index 8f74ab4f09..183ccb5b5a 100644 --- a/src/scvi/module/_vae.py +++ b/src/scvi/module/_vae.py @@ -182,6 +182,7 @@ def __init__( 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 @@ -189,6 +190,7 @@ def __init__( 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: @@ -302,7 +304,6 @@ def _get_inference_input( return { MODULE_KEYS.X_KEY: tensors[REGISTRY_KEYS.X_KEY], MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_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), } @@ -370,7 +371,6 @@ def _regular_inference( if self.use_observed_lib_size: library = torch.log(x.sum(1)).unsqueeze(1) if self.log_variational: - x_ = x_ / x_.mean(1).unsqueeze(1) x_ = torch.log1p(x_) if cont_covs is not None and self.encode_covariates: @@ -382,7 +382,7 @@ def _regular_inference( else: categorical_input = () - if self.encode_covariates and self.batch_representation == "embedding": + 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) @@ -554,7 +554,7 @@ def loss( tensors: dict[str, torch.Tensor], inference_outputs: dict[str, torch.Tensor | Distribution | None], generative_outputs: dict[str, Distribution | None], - kl_weight: float = 1.0, + kl_weight: torch.tensor | float = 1.0, ) -> LossOutput: """Compute the loss.""" from torch.distributions import kl_divergence diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index faf5ebb5a0..daa73ad9dc 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -290,7 +290,7 @@ class Encoder(nn.Module): n_continuous The dimensionality of the continuous covariates including batch embeddings. - n_cat_l)t + n_cat_list A list containing the number of categories for each category of interest. Each category will be included using a one-hot encoding @@ -677,7 +677,7 @@ class MultiEncoder(nn.Module): def __init__( self, n_heads: int, - n_input_list: Iterable[int], + n_input_list: list[int], n_output: int, n_hidden: int = 128, n_layers_individual: int = 1, diff --git a/src/scvi/nn/_embedding.py b/src/scvi/nn/_embedding.py index 956109de2b..2faf885ab7 100644 --- a/src/scvi/nn/_embedding.py +++ b/src/scvi/nn/_embedding.py @@ -19,9 +19,9 @@ def _partial_freeze_hook_factory(freeze: int) -> Callable[[torch.Tensor], torch. """ def _partial_freeze_hook(grad: torch.Tensor) -> torch.Tensor: - grad_copy = grad.clone() - grad_copy[:freeze] = 0.0 - return grad_copy + grad = grad.clone() + grad[:freeze] = 0.0 + return grad return _partial_freeze_hook diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index d47c2ea4d0..2b0e8a3be8 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -20,9 +20,7 @@ from scvi.module.base import ( BaseModuleClass, LossOutput, - MogPrior, PyroBaseModuleClass, - VampPrior, ) from scvi.train._constants import METRIC_KEYS from scvi.utils import is_package_installed @@ -749,15 +747,6 @@ def training_step(self, batch, batch_idx): loss_ = kappa * self.loss_adversarial_classifier( z, adversarial_group, batch_tensor, True ) - if isinstance(self.module.prior, MogPrior) or isinstance( - self.module.prior, VampPrior - ): - qz_m, qz_v = qz.loc.detach(), qz.scale.detach() - loss_ += self.module.prior.kl( - qz=Normal(qz_m, qz_v), - z=z, - labels=batch.get(REGISTRY_KEYS.LABELS_KEY, torch.tensor(0)).long(), - ).mean() if i > 1 and (loss - loss_) / loss < 1e-3: break loss = loss_ diff --git a/src/scvi/utils/_docstrings.py b/src/scvi/utils/_docstrings.py index 84d08e6542..33ad12b0f4 100644 --- a/src/scvi/utils/_docstrings.py +++ b/src/scvi/utils/_docstrings.py @@ -128,7 +128,7 @@ assay_key key in `adata.obs` for assay and suspension type information. Categories will automatically be converted into integer categories and saved to `adata.obs['_scvi_assay']`. If `None`, assigns - the same batch to all the data.""" + the same assay to all the data.""" param_sample_key = """\ sample_key diff --git a/tests/external/muanvi/test_muanvi.py b/tests/external/muanvi/test_muanvi.py deleted file mode 100644 index 1d2cba0e65..0000000000 --- a/tests/external/muanvi/test_muanvi.py +++ /dev/null @@ -1,144 +0,0 @@ -import itertools - -import numpy as np -import pandas as pd -import pytest -from anndata import AnnData -from mudata import MuData -from scipy import sparse as sp_sparse - -from scvi.data import synthetic_iid - - -# helper function for testing purposes ; could be moved in the same file as generate_synthetic() -def _generate_synthetic_hierarchy( - batch_size: int = 128, - n_genes: int = 100, - n_proteins: int = 100, - n_batches: int = 2, - n_labels_1: int = 2, - n_labels_2: int = 11, - n_sites: int = 2, - sparse: bool = False, -) -> AnnData: - """ - New method to generate test data with two-layer labels. - """ - n_total = batch_size * n_batches * n_sites - data = np.random.negative_binomial(5, 0.3, size=(n_total)) - mask = np.random.binomial(n=1, p=0.7, size=(n_total, n_genes)) - data = data * mask # We put the batch index first - labels_1 = np.random.randint(0, n_labels_1, size=(n_total,)) - labels_1 = np.array([f"label_{i}" for i in labels_1]) - labels_2 = np.random.randint(0, n_labels_2 // n_sites, size=(n_total,)) - - batch = [] - site = [] - labels_2 = [] - for site, batch in itertools.product(range(n_sites), range(n_batches)): - batch += [f"batch_{batch}_site_{site}"] * batch_size - site += [f"site_{site}"] * batch_size - labels_1 = np.array([f"label_{d}_{s}" for d, s in zip(labels_2, site, strict=True)]) - - if sparse: - data = sp_sparse.csr_matrix(data) - adata = AnnData(data) - adata.obs["batch"] = pd.Categorical(batch) - adata.obs["site"] = pd.Categorical(site) - adata.obs["labels_1"] = pd.Categorical(labels_1) - adata.obs["labels_2"] = pd.Categorical(labels_2) - - # Protein measurements - p_data = np.random.negative_binomial(5, 0.3, size=(adata.shape[0], n_proteins)) - adata.obsm["protein_expression"] = p_data - adata.uns["protein_names"] = np.arange(n_proteins).astype(str) - - return adata - - -# helper function for testing purposes ; could be moved int he same file as synthetic_iid() -def synthetic_iid_hierarchy( - batch_size: int | None = 200, - n_genes: int | None = 100, - n_proteins: int | None = 100, - n_batches: int | None = 2, - n_sites: int | None = 2, - n_labels_1: int | None = 2, - n_labels_2: int | None = 10, - sparse: bool = False, -) -> AnnData: - """Synthetic dataset with ZINB distributed RNA and NB distributed protein, with three-layer annotation. - This dataset is just for testing purposed and not meant for modeling or research. - Each value is independently and identically distributed. - Parameters - ---------- - batch_size - Number of cells per batch - n_genes - Number of genes - n_proteins - Number of proteins - n_batches - Number of batches - n_sites - Number of sites - n_labels_1 - Number of cell types of layer 1 - n_labels_2 - Number of cell types of layer 2 - sparse - Whether to use a sparse matrix - Returns - ------- - AnnData with batch info (``.obs['batch']``), label info (``.obs['labels']``), - site info (``.obs['site']``), - protein expression (``.obsm["protein_expression"]``) and - protein names (``.obs['protein_names']``) - Examples - -------- - >>> import scvi - >>> adata = scvi.data.synthetic_iid() - """ - - return _generate_synthetic_hierarchy( - batch_size=batch_size, - n_genes=n_genes, - n_proteins=n_proteins, - n_batches=n_batches, - n_sites=n_sites, - n_labels_1=n_labels_1, - n_labels_2=n_labels_2, - sparse=sparse, - ) - - -def test_methylvi(): - adata1 = synthetic_iid() - adata1.layers["mc"] = adata1.X - adata1.layers["cov"] = adata1.layers["mc"] + 10 - - adata2 = synthetic_iid() - adata2.layers["mc"] = adata2.X - adata2.layers["cov"] = adata2.layers["mc"] + 10 - - mdata = MuData({"mod1": adata1, "mod2": adata2}) - - METHYLVI.setup_mudata( - mdata, - mc_layer="mc", - cov_layer="cov", - methylation_contexts=["mod1", "mod2"], - batch_key="batch", - modalities={"batch_key": "mod1"}, - ) - vae = METHYLVI( - mdata, - ) - vae.train(3) - vae.get_elbo(indices=vae.validation_indices) - vae.get_normalized_methylation() # Retrieve methylation for all contexts - vae.get_normalized_methylation(context="mod1") # Retrieve for specific context - with pytest.raises(ValueError): # Should fail when invalid context selected - vae.get_normalized_methylation(context="mod3") - vae.get_latent_representation() - vae.differential_methylation(groupby="mod1:labels", group1="label_1") From 6027d9dba3d7abe3c0626bba637b545ae006127c Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 22 May 2026 09:35:19 +0000 Subject: [PATCH 17/24] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/scvi/train/_trainingplans.py | 1 - 1 file changed, 1 deletion(-) diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 2b0e8a3be8..d9c8f49dde 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -12,7 +12,6 @@ import torchmetrics.functional as tmf from lightning.pytorch.strategies.ddp import DDPStrategy from pyro.nn import PyroModule -from torch.distributions import Normal from torch.optim.lr_scheduler import ReduceLROnPlateau from scvi import REGISTRY_KEYS, settings From 45d6a162ea84e845d34d90e95a0e8836b8531154 Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Sat, 23 May 2026 00:43:16 +0200 Subject: [PATCH 18/24] Fixes tests. --- src/scvi/external/sysvi/_base_components.py | 2 +- src/scvi/external/sysvi/_module.py | 4 ++-- src/scvi/external/sysvi/_priors.py | 2 +- src/scvi/module/base/_embedding_mixin.py | 2 +- src/scvi/nn/_base_components.py | 13 ++++++------- src/scvi/train/_trainingplans.py | 17 +++++++++++------ 6 files changed, 22 insertions(+), 18 deletions(-) diff --git a/src/scvi/external/sysvi/_base_components.py b/src/scvi/external/sysvi/_base_components.py index 98430d3beb..375f9dbbb4 100644 --- a/src/scvi/external/sysvi/_base_components.py +++ b/src/scvi/external/sysvi/_base_components.py @@ -116,7 +116,7 @@ def forward( parametrized with the predicted parameters. """ cat_list = [batch_index] + cat_list - q_ = self.decoder_y(x, *cat_list, cont=cont) + q_ = self.decoder_y(x, *cat_list, cont_input=cont) q_m = self.mean_encoder(q_) if q_m.isnan().any() or q_m.isinf().any(): warnings.warn( diff --git a/src/scvi/external/sysvi/_module.py b/src/scvi/external/sysvi/_module.py index 9712845476..60086ba40e 100644 --- a/src/scvi/external/sysvi/_module.py +++ b/src/scvi/external/sysvi/_module.py @@ -9,7 +9,7 @@ from scvi.module.base import BaseModuleClass, EmbeddingModuleMixin, LossOutput, auto_move_data from ._base_components import EncoderDecoder -from ._priors import GaussianPrior, VampPrior +from ._priors import StandardPrior, VampPrior if TYPE_CHECKING: from typing import Literal @@ -137,7 +137,7 @@ def __init__( ) if prior == "standard_normal": - self.prior = GaussianPrior() + self.prior = StandardPrior() elif prior == "vamp": assert pseudoinput_data is not None, ( "Pseudoinput data must be specified if using VampPrior" diff --git a/src/scvi/external/sysvi/_priors.py b/src/scvi/external/sysvi/_priors.py index 89d615f64e..c966fafb9e 100644 --- a/src/scvi/external/sysvi/_priors.py +++ b/src/scvi/external/sysvi/_priors.py @@ -35,7 +35,7 @@ def kl( pass -class GaussianPrior(Prior): +class StandardPrior(Prior): """Standard prior distribution.""" def kl(self, qz: torch.Tensor, z: None = None) -> torch.Tensor: diff --git a/src/scvi/module/base/_embedding_mixin.py b/src/scvi/module/base/_embedding_mixin.py index 02dd4ebef9..3a58b7d868 100644 --- a/src/scvi/module/base/_embedding_mixin.py +++ b/src/scvi/module/base/_embedding_mixin.py @@ -19,7 +19,7 @@ def embeddings_dict(self) -> ModuleDict: @property def embeddings_dim(self) -> dict: """Dictionary of embeddings dimensions.""" - if not hasattr(self, "_embeddings_dict"): + if not hasattr(self, "_embeddings_dim"): self._embeddings_dim = {} return self._embeddings_dim diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index daa73ad9dc..0e6efc6a27 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -134,7 +134,7 @@ def __init__( f"Layer {i}", nn.Sequential( nn.Linear( - n_in + (cat_dim + n_continuous) * self.inject_into_layer(i), + n_in + self.n_cov * self.inject_into_layer(i), n_out, bias=bias, ), @@ -213,12 +213,12 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No tensor of shape ``(n_out,)`` """ one_hot_cat_list = [] # for generality in this list many idxs useless. - cont_list = [cont] if cont is not None else [] + cont_list = [cont_input] if cont_input is not None else [] cat_list = cat_list or [] if len(self.n_cat_list) > len(cat_list): raise ValueError("nb. categorical args provided doesn't match init. params.") - if self.n_continuous > 0 and cont_input.shape[-1] != self.n_continuous: + if self.n_continuous > 0 and cont_input is not None and cont_input.shape[-1] != self.n_continuous: raise ValueError("continuous dims provided doesn't match init. params.") for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): if n_cat and cat is None: @@ -229,8 +229,7 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No else: one_hot_cat = cat # cat has already been one_hot encoded one_hot_cat_list += [one_hot_cat] - if cont_input is not None: - one_hot_cat_list += [cont_input] + cov_list = cont_list + one_hot_cat_list for i, layers in enumerate(self.fc_layers): for layer in layers: if layer is not None: @@ -382,7 +381,7 @@ def forward( """ # Parameters for latent distribution - q = self.encoder(x, *cat_list, cont=cont) + q = self.encoder(x, *cat_list, cont_input=cont) q_m = self.mean_encoder(q) q_v = self.var_activation(self.var_encoder(q)) + self.var_eps dist = Normal(q_m, q_v.sqrt()) @@ -521,7 +520,7 @@ def forward( """ # The decoder returns values for the parameters of the ZINB distribution - px = self.px_decoder(z, *cat_list, cont=cont) + px = self.px_decoder(z, *cat_list, cont_input=cont) if output_condition is not None and self.n_conditions_output: one_hot_cat = nn.functional.one_hot( output_condition.squeeze(-1), self.n_conditions_output diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index d9c8f49dde..84002386d1 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -651,7 +651,7 @@ def __init__( self.adversarial_classifier = False else: self.adversarial_classifier = Classifier( - n_input=self.module.n_latent + self.module.n_adversarial_group, + n_input=self.module.n_latent + getattr(self.module, "n_adversarial_group", 0), n_hidden=128, n_labels=self.n_output_classifier, n_layers=2, @@ -670,12 +670,17 @@ def loss_adversarial_classifier( ): """Loss for adversarial classifier.""" n_classes = self.n_output_classifier - adversarial_group_ = torch.nn.functional.one_hot( - adversarial_group, num_classes=self.module.n_adversarial_group - ).float() + n_adv_group = getattr(self.module, "n_adversarial_group", 0) + if n_adv_group > 0: + adversarial_group_ = torch.nn.functional.one_hot( + adversarial_group, num_classes=n_adv_group + ).float() + z_cls = torch.cat([z, adversarial_group_], dim=1) + else: + z_cls = z if predict_true_class: # train classifier - z = z.detach() - z = torch.cat([z, adversarial_group_], dim=1) + z_cls = z_cls.detach() + z = z_cls cls_logits = self.adversarial_classifier(z) if predict_true_class: From b90f88a9dbbf44fe22e9308da48a06f959f4be0b Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 22 May 2026 22:45:30 +0000 Subject: [PATCH 19/24] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/scvi/nn/_base_components.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index 0e6efc6a27..fde19f3c2f 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -218,7 +218,11 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No if len(self.n_cat_list) > len(cat_list): raise ValueError("nb. categorical args provided doesn't match init. params.") - if self.n_continuous > 0 and cont_input is not None and cont_input.shape[-1] != self.n_continuous: + if ( + self.n_continuous > 0 + and cont_input is not None + and cont_input.shape[-1] != self.n_continuous + ): raise ValueError("continuous dims provided doesn't match init. params.") for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): if n_cat and cat is None: From 3820a338c1a415264359424e25568451de943b0d Mon Sep 17 00:00:00 2001 From: Can Ergen Date: Sat, 23 May 2026 23:24:45 +0200 Subject: [PATCH 20/24] Small fixes. --- docs/user_guide/models/scvix.md | 268 +++++++++++++++++++++++++++++ src/scvi/external/scvix/_model.py | 147 +++++++++++++--- src/scvi/external/scvix/_module.py | 17 -- tests/external/scvix/test_scvix.py | 35 ++++ 4 files changed, 423 insertions(+), 44 deletions(-) create mode 100644 docs/user_guide/models/scvix.md diff --git a/docs/user_guide/models/scvix.md b/docs/user_guide/models/scvix.md new file mode 100644 index 0000000000..cfbfb4f972 --- /dev/null +++ b/docs/user_guide/models/scvix.md @@ -0,0 +1,268 @@ +# scVI-X + +**scVI-X** [^ref1] (single-cell Variational Inference — cross-technology; +Python class {class}`~scvi.external.SCVIX`) is a deep generative model for +learning technology-invariant cell-state representations across diverse +single-cell sequencing platforms. It extends scVI [^ref2] with three targeted +architectural changes that together enable strong cross-technology integration +and zero-shot query mapping without fine-tuning. + +The key advantages of scVI-X over scVI are: + +- **Cross-technology integration**: Explicitly models the sequencing assay + (e.g., 10x single-cell vs. single-nucleus) as a distinct source of + variation, yielding substantially better mixing across platforms. +- **Zero-shot query mapping**: Trained models can embed new cells from + previously unseen arrays via a single forward pass, with no retraining + required. +- **Assay-level adversarial training**: Adversarial correction is applied at + the assay level rather than the batch level, reducing the risk of + overintegration when cell-type compositions differ between batches. +- **Assay-aware decoding**: Learned gene-level assay biases in the output + layer capture technology-specific transcript-detection differences and + provide interpretable, reproducible signatures across datasets. +- **Variational batch embeddings**: Optionally replaces one-hot batch + encodings with low-dimensional variational embeddings, reducing parameter + count and improving identifiability through a KL regularisation term. + +```{topic} Tutorials: + +- {doc}`/tutorials/notebooks/scrna/scvix_tutorial` +``` + +## Preliminaries + +scVI-X takes as input a scRNA-seq count matrix $X$ with $N$ cells and $G$ +genes. For each cell $n$ it requires: + +- An **assay label** $a_n$ identifying the sequencing technology or + suspension type (e.g., 10x Chromium single-cell, 10x Chromium + single-nucleus, Smart-seq2). This is the primary covariate that scVI-X is + designed to correct for. +- A **batch label** $b_n$ for experimental nuisance variation within an + assay (e.g., sequencing run, donor, processing day). + +Optionally, scVI-X can also accept: + +- Cell-type labels $y_n$ for supervised integration via a cell-type + biased prior. +- An **adversarial group label** $g_n$ (defaults to $a_n$) that defines + which grouping the adversarial classifier acts on. To reduce + overintegration if cell-type composition varies across assays. +- Additional categorical or continuous covariates. + +## Generative process + +scVI-X posits a conditional VAE generative model. For each cell $n$: + +1. A latent cell-state vector $z_n$ is drawn from the prior: + + $z_n \sim p(z)$ + + where $p(z)$ is one of: a standard Gaussian, a mixture of Gaussians + (MOG), a variational amortized mixture of posteriors (VampPrior). + +2. The normalised expression scale $h_n$ is decoded from $z_n$ conditioned + on the batch $b_n$ and assay $a_n$: + + $h_n = \mathrm{softmax}\!\left(f_\theta(z_n,\, b_n) + \delta_{a_n}\right)$ + + where $f_\theta$ is a fully-connected decoder network, and + $\delta_{a_n} \in \mathbb{R}^G$ is a learned, assay-specific output bias + vector that captures gene-level differences in transcript-detection + efficiency across sequencing platforms. + +3. Observed counts are generated from a negative binomial (or ZINB / Poisson) + likelihood: + + $x_{ng} \mid h_{ng} \sim \mathrm{NegativeBinomial}(l_n\, h_{ng},\; r_{ng})$ + + where $l_n$ is the observed library size and $r_{ng}$ is the + gene-specific (or cell-specific) inverse dispersion. + +The latent variables are summarised below: + +```{eval-rst} +.. list-table:: + :widths: 20 80 15 + :header-rows: 1 + + * - Variable + - Description + - Code name + * - :math:`z_n \in \mathbb{R}^L` + - Technology-invariant cell-state embedding. + - ``z`` + * - :math:`h_n \in \mathbb{R}^G` + - Normalised gene expression scale. + - ``px_scale`` + * - :math:`l_n \in \mathbb{R}^+` + - Observed library size (log-transformed). + - ``library`` + * - :math:`r_{ng} \in \mathbb{R}^+` + - Gene- (and optionally cell-) specific inverse dispersion. + - ``px_r`` + * - :math:`\delta_{a_n} \in \mathbb{R}^G` + - Assay-specific gene-level output bias. + - output layer weights conditioned on ``assay_index`` +``` + +### Variational batch embeddings (optional) + +When `batch_representation="embedding"` is selected, each batch $b$ is +represented by a low-dimensional embedding +$e_b \sim \mathcal{N}(\mu_b, \sigma_b^2 I)$ rather than a one-hot vector. +This reduces the number of decoder parameters. This embedding can be variational +and then adds a KL regularisation +term to the loss that improves identifiability of batch representations: + +$\mathcal{L}_\text{embed} = \mathrm{KL}\!\left(q(e_b) \,\|\, \mathcal{N}(0,I)\right)$ + +## Inference + +scVI-X uses amortised variational inference. The approximate posterior is: + +$q_\phi(z_n \mid x_n, a_n) = \mathcal{N}\!\left(\mu_\phi(x_n, a_n),\; \sigma^2_\phi(x_n, a_n)\, I\right)$ + +where $\mu_\phi$ and $\sigma^2_\phi$ are encoder neural networks. The assay +index $a_n$ is injected into the encoder through **conditional layer +normalisation** (controlled by `conditional_norm=True`, the default): instead +of a single global scale and shift applied after each hidden layer, the model +learns per-assay scale and shift parameters $\gamma_{a_n}$ and $\beta_{a_n}$. +This provides a lightweight but effective mechanism for correcting +assay-specific distributional shifts in the encoder. + +The evidence lower bound (ELBO) trained is: + +$\mathcal{L} = \mathbb{E}_{q_\phi(z_n \mid x_n, a_n)}\!\left[\log p_\theta(x_n \mid z_n, b_n, a_n)\right] - \mathrm{KL}\!\left(q_\phi(z_n \mid x_n, a_n) \,\|\, p(z_n)\right) + \mathcal{L}_\text{embed}$ + +### Adversarial training + +When multiple assays are present, scVI-X adds an adversarial +classifier that acts on the latent space $z_n$. The classifier is trained to +predict the adversarial label $a_n$ (typically the assay), while the +encoder is simultaneously trained to fool it: + +$\mathcal{L}_\text{adv} = -\mathbb{E}\left[\log p_\psi(\hat{a}_n \neq a_n \mid z_n)\right]$ + +By conditioning adversarial correction at the **assay** level rather than the +fine-grained batch level, scVI-X reduces the risk of overintegration in +settings where cell-type compositions differ between donors or conditions +within the same assay. + +Adversarial training is automatically enabled when more than one assay is +present; the adversarial label can be set to any registered categorical +covariate via `adversarial_key`. + +## Prior options + +scVI-X supports four prior distributions for $z_n$, selected via the `prior` +argument at model initialisation. We found Gaussian to perform well in all tested scenarios. + +- **`'gaussian'`** (default): Standard normal $\mathcal{N}(0, I)$. Fastest + to train; suitable for most use cases. +- **`'mog'`**: Mixture of Gaussians with $K$ learnable components. Encourages + a more structured latent space. +- **`'vamp'`**: VampPrior ([Tomczak & Welling, 2018](https://doi.org/10.48550/arXiv.1705.07120)). + Prior modes are anchored to learned pseudoinputs drawn from the data, + making the prior data-adaptive. +- **`'mog_celltype'`**: Mixture of Gaussians with one component per cell-type + label. Requires `labels_key` to be set in `setup_anndata`. Guides + integration by biasing the prior toward annotated cell-type clusters. + +## Tasks + +### Latent representation + +The primary output of scVI-X is a low-dimensional, technology-corrected +embedding of each cell: + +```python +import scvi + +scvi.external.SCVIX.setup_anndata( + adata, + batch_key="batch", + assay_key="assay", # sequencing technology / suspension type +) +model = scvi.external.SCVIX(adata) # defaults: embedding, n_latent=20, n_layers=2, n_hidden=512 +model.train() # defaults: batch_size=1024, n_epochs_kl_warmup=5 +adata.obsm["X_scVIX"] = model.get_latent_representation() +``` + +### Normalised expression + +Batch- and assay-corrected normalised expression can be obtained by decoding +from the latent space while fixing the batch and/or assay to a reference +value. This is analogous to `get_normalized_expression` in scVI and supports +the `transform_batch` and `transform_assay` arguments: + +```python +# Expression as it would look in a given reference batch +norm_expr = model.get_normalized_expression(transform_batch="batch_1") +``` + +### Zero-shot query mapping + +scVI-X supports query-to-reference mapping via the scArches framework without +any retraining. A query dataset with new batches (but the same gene set and assay) can +be embedded directly without additional training. + +```python +# Save the reference model +model.save("scvix_reference/") + +# Prepare the query and embed zero-shot +scvi.external.SCVIX.prepare_query_anndata(adata_query, "scvix_reference/") +query_model = scvi.external.SCVIX.load_query_data(adata_query, "scvix_reference/") +query_model.train(max_epochs=200, plan_kwargs={"weight_decay": 0.0}) +adata_query.obsm["X_scVIX"] = query_model.get_latent_representation() +``` + +### Differential expression + +Standard differential expression analysis between cell groups is available +via the inherited `differential_expression` method: + +```python +de_results = model.differential_expression(groupby="cell_type", group1="T cell") +``` + +### Batch embeddings + +When `batch_representation="embedding"` is used, the learned batch embeddings +can be retrieved and used for downstream analysis (e.g., visualising +batch-level variation or clustering batches): + +```python +model = scvi.external.SCVIX(adata, batch_representation="embedding") +model.train() +batch_emb = model.get_batch_representation() # shape: (n_cells, embedding_dim) +``` + +## Key setup_anndata arguments + +```{eval-rst} +.. list-table:: + :widths: 25 75 + :header-rows: 1 + + * - Argument + - Description + * - ``batch_key`` + - Column in ``adata.obs`` for experimental batch (donor, run, etc.). + * - ``assay_key`` + - Column in ``adata.obs`` for sequencing technology / suspension type. + Drives conditional layer normalisation and assay-output biases. +``` + +[^ref1]: + Can Ergen, Ori Kronfeld, Martin Kim, Shiyi Yang, Florian Ingelfinger, + Nir Yosef (in preparation), + _scvi-X: Learning technology-invariant cell states_. + +[^ref2]: + Romain Lopez, Jeffrey Regier, Michael B Cole, Michael I Jordan, Nir Yosef + (2018), + _Deep generative modeling for single-cell transcriptomics_, + [Nature Methods](https://doi.org/10.1038/s41592-018-0229-2). diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index c1247aa1f0..a4b094eabd 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -36,6 +36,7 @@ from typing import Literal import numpy as np + import pandas as pd from anndata import AnnData logger = logging.getLogger(__name__) @@ -128,19 +129,24 @@ class SCVIX( def __init__( self, adata: AnnData | None = None, - n_hidden: int = 128, - n_latent: int = 10, - n_layers: int = 1, + n_hidden: int = 512, + n_latent: int = 20, + n_layers: int = 2, dropout_rate: float = 0.05, - dispersion: Literal["gene", "gene-batch", "gene-cell"] = "gene", + dispersion: Literal["gene", "gene-batch", "gene-assay", "gene-cell"] = "gene-assay", gene_likelihood: Literal["zinb", "nb", "poisson", "normal"] = "nb", prior: Literal["gaussian", "mog", "vamp", "mog_celltype"] = "gaussian", + batch_representation: Literal["one-hot", "embedding"] = "embedding", + batch_embedding_kwargs: dict | None = None, pseudoinputs_data_indices: np.array | None = None, n_prior_components: int = 50, **kwargs, ): super().__init__(adata) + if batch_embedding_kwargs is None: + batch_embedding_kwargs = {"variational": True, "embedding_dim": n_latent} + self._module_kwargs = { "n_hidden": n_hidden, "n_latent": n_latent, @@ -148,13 +154,16 @@ def __init__( "dropout_rate": dropout_rate, "dispersion": dispersion, "gene_likelihood": gene_likelihood, + "batch_representation": batch_representation, + "batch_embedding_kwargs": batch_embedding_kwargs, **kwargs, } self._model_summary_string = ( "SCVI model with the following parameters: \n" f"n_hidden: {n_hidden}, n_latent: {n_latent}, n_layers: {n_layers}, " f"dropout_rate: {dropout_rate}, dispersion: {dispersion}, " - f"gene_likelihood: {gene_likelihood}." + f"gene_likelihood: {gene_likelihood}, " + f"batch_representation: {batch_representation}." ) if prior == "vamp": @@ -196,6 +205,8 @@ def __init__( dropout_rate=dropout_rate, dispersion=dispersion, gene_likelihood=gene_likelihood, + batch_representation=batch_representation, + batch_embedding_kwargs=batch_embedding_kwargs, prior=prior, pseudoinput_data=pseudoinput_data, n_prior_components=n_prior_components, @@ -215,14 +226,14 @@ def train( train_size: float | None = None, validation_size: float | None = None, shuffle_set_split: bool = True, - batch_size: int = 256, - early_stopping: bool = True, - check_val_every_n_epoch: int | None = None, - reduce_lr_on_plateau: bool = True, - n_steps_kl_warmup: int | None = None, - n_epochs_kl_warmup: int | None = None, + batch_size: int = 1024, + early_stopping: bool = False, + n_epochs_kl_warmup: int | None = 5, adversarial_classifier: bool | None = None, adversarial_key: str = "assay", + scale_adversarial_loss: float | str = 5.0, + adversarial_steps: int = 5, + weight_kl_sample: float = 1e-5, datasplitter_kwargs: dict | None = None, plan_kwargs: dict | None = None, external_indexing: list[np.array] = None, @@ -251,20 +262,18 @@ def train( Minibatch size to use during training. early_stopping Whether to perform early stopping with respect to the validation set. - check_val_every_n_epoch - Check val every n train epochs. By default, val is not checked, unless `early_stopping` - is `True` or `reduce_lr_on_plateau` is `True`. If either of the latter conditions are - met, val is checked every epoch. - reduce_lr_on_plateau - Reduce learning rate on plateau of validation metric (default is ELBO). n_epochs_kl_warmup Number of epochs to scale weight on KL divergences from 0 to 1. - scale_adversarial_classifier - How to weight adversarial classifier in the latent space. This helps mixing when - there are multiple assays. Defaults to `1`. adversarial_key - Key in `adata.obs` that corresponds to batch or assay key to use for adversarial - training. If `None`, defaults to the assay key. + Key in `adata.obs` that corresponds to the batch or assay key to use for adversarial + training. Defaults to the assay key. + scale_adversarial_loss + Scaling factor for the adversarial loss component. Higher values enforce stronger + assay mixing. Use ``"auto"`` to scale automatically with the KL warmup weight. + adversarial_steps + Number of adversarial classifier gradient steps per training step. + weight_kl_sample + Weight on the KL divergence of the variational batch embeddings. datasplitter_kwargs Additional keyword arguments passed into :class:`~scvi.dataloaders.DataSplitter`. plan_kwargs @@ -284,16 +293,15 @@ def train( n_epochs_kl_warmup = ( n_epochs_kl_warmup if n_epochs_kl_warmup is not None else max_epochs // 2 ) - if reduce_lr_on_plateau: - check_val_every_n_epoch = 1 update_dict = { "lr": lr, "adversarial_classifier": adversarial_classifier, "adversarial_key": adversarial_key, - "reduce_lr_on_plateau": reduce_lr_on_plateau, "n_epochs_kl_warmup": n_epochs_kl_warmup, - "n_steps_kl_warmup": n_steps_kl_warmup, + "scale_adversarial_loss": scale_adversarial_loss, + "adversarial_steps": adversarial_steps, + "weight_kl_sample": weight_kl_sample, } if plan_kwargs is not None: plan_kwargs.update(update_dict) @@ -324,11 +332,96 @@ def train( accelerator=accelerator, devices=devices, early_stopping=early_stopping, - check_val_every_n_epoch=check_val_every_n_epoch, **kwargs, ) return runner() + def _get_transform_batch_gen_kwargs(self, batch): + kwargs = super()._get_transform_batch_gen_kwargs(batch) + if hasattr(self, "_transform_assay_override"): + kwargs["transform_assay"] = self._transform_assay_override + return kwargs + + def get_normalized_expression( + self, + adata: AnnData | None = None, + indices: list[int] | None = None, + transform_batch: list[int | str] | None = None, + transform_assay: int | str | None = None, + gene_list: list[str] | None = None, + library_size: float | Literal["latent"] = 1, + n_samples: int = 1, + n_samples_overall: int | None = None, + weights: Literal["uniform", "importance"] | None = None, + batch_size: int | None = None, + return_mean: bool = True, + return_numpy: bool | None = None, + silent: bool = True, + dataloader=None, + data_loader_kwargs: dict | None = None, + **importance_weighting_kwargs, + ) -> np.ndarray | pd.DataFrame: + """Returns the normalized (decoded) gene expression. + + Extends the base ``get_normalized_expression`` with assay conditioning via + ``transform_assay``. + + Parameters + ---------- + transform_assay + Assay to condition on for all cells when decoding. If a string, must be a valid + assay category registered via :meth:`~scvi.external.SCVIX.setup_anndata`. If an + integer, used directly as the assay index. If ``None`` (default), each cell is + decoded using its own observed assay. + transform_batch + Batch to condition on. See base class for details. + + Notes + ----- + All other parameters are identical to + :meth:`~scvi.model.base.RNASeqMixin.get_normalized_expression`. + """ + if transform_assay is not None: + adata_ = self._validate_anndata(adata) + assay_mappings = ( + self.get_anndata_manager(adata_, required=True) + .get_state_registry(REGISTRY_KEYS.ASSAY_KEY) + .categorical_mapping + ) + # Hack for transform_assay to be passed in get_normalized_expression. + if isinstance(transform_assay, str): + if transform_assay not in assay_mappings: + raise ValueError(f'"{transform_assay}" is not a valid assay category.') + self._transform_assay_override = int( + np.where(assay_mappings == transform_assay)[0][0] + ) + else: + self._transform_assay_override = int(transform_assay) + + try: + result = super().get_normalized_expression( + adata=adata, + indices=indices, + transform_batch=transform_batch, + gene_list=gene_list, + library_size=library_size, + n_samples=n_samples, + n_samples_overall=n_samples_overall, + weights=weights, + batch_size=batch_size, + return_mean=return_mean, + return_numpy=return_numpy, + silent=silent, + dataloader=dataloader, + data_loader_kwargs=data_loader_kwargs, + **importance_weighting_kwargs, + ) + finally: + if hasattr(self, "_transform_assay_override"): + del self._transform_assay_override + + return result + @classmethod @setup_anndata_dsp.dedent def setup_anndata( diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index a314a09c21..fa17862904 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -533,7 +533,6 @@ def loss( generative_outputs: dict[str, Distribution | None], kl_weight: float = 1.0, weight_assay_loss: float = 0.0, - weight_global: float = 1.0, weight_kl_sample: float = 1.0, classification_ratio: float = 500.0, ) -> LossOutput: @@ -570,28 +569,12 @@ def loss( reconst_loss = -generative_outputs[MODULE_KEYS.PX_KEY].log_prob(x).sum(-1) - kl_global = torch.zeros_like(kl_divergence_z) - - if self.gene_likelihood == "zinb": - zi = generative_outputs[MODULE_KEYS.PX_KEY].zi_logits - kl_global -= ( - distributions.Exponential(10.0 * torch.ones_like(zi)) - .log_prob(torch.exp(zi)) - .sum(-1) - ) - if self.gene_likelihood == "zinb" or self.gene_likelihood == "nb": - theta = generative_outputs[MODULE_KEYS.PX_KEY].theta - kl_global -= ( - distributions.Exponential(torch.ones_like(theta)).log_prob(1 / theta).sum(-1) - ) - weighted_kl_local = kl_weight * (kl_divergence_z + weight_kl_sample * kl_divergence_sample) loss = torch.mean( reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss - + weight_global * kl_global ) if self.n_labels > 1: diff --git a/tests/external/scvix/test_scvix.py b/tests/external/scvix/test_scvix.py index 0d926ffce1..bd6e47f915 100644 --- a/tests/external/scvix/test_scvix.py +++ b/tests/external/scvix/test_scvix.py @@ -41,6 +41,41 @@ def test_scvix_dispersion(dispersion: str): model.get_normalized_expression() +def test_scvix_transform_assay(): + adata = synthetic_iid(batch_size=100) + # Use batch as both batch and assay (two categories: batch_0, batch_1) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch") + model = SCVIX(adata) + model.train(max_epochs=1) + + # baseline: no transformation + expr_default = model.get_normalized_expression(return_numpy=True) + + # transform by integer index (assay 0) + scvi.settings.seed = 42 + expr_int = model.get_normalized_expression(transform_assay=0, return_numpy=True) + assert expr_int.shape == expr_default.shape + + # transform by string name — must match integer result (same seed, same code) + scvi.settings.seed = 42 + expr_str = model.get_normalized_expression(transform_assay="batch_0", return_numpy=True) + np.testing.assert_array_equal(expr_int, expr_str) + + # transform_assay and transform_batch can be combined + expr_both = model.get_normalized_expression( + transform_assay="batch_0", transform_batch="batch_1", return_numpy=True + ) + assert expr_both.shape == expr_default.shape + + # different assay indices produce different outputs + expr_assay1 = model.get_normalized_expression(transform_assay=1, return_numpy=True) + assert not np.allclose(expr_int, expr_assay1) + + # invalid assay name raises ValueError + with pytest.raises(ValueError, match="not a valid assay category"): + model.get_normalized_expression(transform_assay="nonexistent_assay") + + def test_scvix_encode_covariates(): adata = synthetic_iid(batch_size=100) SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch", labels_key="labels") From c5e06c6758d33ab7940848ecbe6a93d31cb3f273 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sat, 23 May 2026 21:24:59 +0000 Subject: [PATCH 21/24] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- docs/user_guide/models/scvix.md | 10 ++++++---- src/scvi/external/scvix/_module.py | 6 +----- 2 files changed, 7 insertions(+), 9 deletions(-) diff --git a/docs/user_guide/models/scvix.md b/docs/user_guide/models/scvix.md index cfbfb4f972..48ba92664a 100644 --- a/docs/user_guide/models/scvix.md +++ b/docs/user_guide/models/scvix.md @@ -183,10 +183,12 @@ import scvi scvi.external.SCVIX.setup_anndata( adata, batch_key="batch", - assay_key="assay", # sequencing technology / suspension type + assay_key="assay", # sequencing technology / suspension type ) -model = scvi.external.SCVIX(adata) # defaults: embedding, n_latent=20, n_layers=2, n_hidden=512 -model.train() # defaults: batch_size=1024, n_epochs_kl_warmup=5 +model = scvi.external.SCVIX( + adata +) # defaults: embedding, n_latent=20, n_layers=2, n_hidden=512 +model.train() # defaults: batch_size=1024, n_epochs_kl_warmup=5 adata.obsm["X_scVIX"] = model.get_latent_representation() ``` @@ -237,7 +239,7 @@ batch-level variation or clustering batches): ```python model = scvi.external.SCVIX(adata, batch_representation="embedding") model.train() -batch_emb = model.get_batch_representation() # shape: (n_cells, embedding_dim) +batch_emb = model.get_batch_representation() # shape: (n_cells, embedding_dim) ``` ## Key setup_anndata arguments diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index fa17862904..94827ddaad 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -571,11 +571,7 @@ def loss( weighted_kl_local = kl_weight * (kl_divergence_z + weight_kl_sample * kl_divergence_sample) - loss = torch.mean( - reconst_loss - + weighted_kl_local - + weight_assay_loss * assay_loss - ) + loss = torch.mean(reconst_loss + weighted_kl_local + weight_assay_loss * assay_loss) if self.n_labels > 1: logits = self.classifier(inference_outputs[MODULE_KEYS.Z_KEY]) From b2e6edf3a0c78d8c4dfc5ef52ae5231a8c704339 Mon Sep 17 00:00:00 2001 From: ori-kron-wis Date: Sun, 24 May 2026 16:23:02 +0300 Subject: [PATCH 22/24] fixing errors from github workflow, pre-commit issues --- src/scvi/external/gimvi/_task.py | 10 ++++++++-- src/scvi/external/scvix/_model.py | 2 -- src/scvi/module/_multivae.py | 4 +++- src/scvi/nn/_base_components.py | 3 ++- 4 files changed, 13 insertions(+), 6 deletions(-) diff --git a/src/scvi/external/gimvi/_task.py b/src/scvi/external/gimvi/_task.py index 04aa97702d..d82e4ab9e2 100644 --- a/src/scvi/external/gimvi/_task.py +++ b/src/scvi/external/gimvi/_task.py @@ -66,7 +66,10 @@ def training_step(self, batch, batch_idx): ] if kappa > 0 and self.adversarial_classifier is not False: fool_loss = self.loss_adversarial_classifier( - torch.cat(zs), torch.cat(batch_tensor).long(), False + torch.cat(zs), + torch.cat(batch_tensor).long(), + torch.cat(batch_tensor).long(), + predict_true_class=False, ) loss += fool_loss * kappa opt1.zero_grad() @@ -94,7 +97,10 @@ def training_step(self, batch, batch_idx): torch.zeros((z.shape[0], 1), device=z.device) + i for i, z in enumerate(zs) ] loss = self.loss_adversarial_classifier( - torch.cat(zs).detach(), torch.cat(batch_tensor).long(), True + torch.cat(zs).detach(), + torch.cat(batch_tensor).long(), + torch.cat(batch_tensor).long(), + predict_true_class=True, ) loss *= kappa opt2.zero_grad() diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index a4b094eabd..510c545179 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -27,7 +27,6 @@ VAEMixin, ) from scvi.train import AdversarialTrainingPlan, TrainRunner -from scvi.utils import setup_anndata_dsp from scvi.utils._docstrings import devices_dsp, setup_anndata_dsp from ._module import VAEX @@ -35,7 +34,6 @@ if TYPE_CHECKING: from typing import Literal - import numpy as np import pandas as pd from anndata import AnnData diff --git a/src/scvi/module/_multivae.py b/src/scvi/module/_multivae.py index 58e4efdab9..f9ea1c1fdb 100644 --- a/src/scvi/module/_multivae.py +++ b/src/scvi/module/_multivae.py @@ -673,12 +673,14 @@ def unsqz(zt, n_s): libsize_acc = unsqz(libsize_acc, n_samples) # sample from the mixed representation - untran_z = Normal(qz_m, qz_v.sqrt()).rsample() + qz = Normal(qz_m, qz_v.sqrt()) + untran_z = qz.rsample() z = self.z_encoder_accessibility.z_transformation(untran_z) outputs = { "x": x, "z": z, + "qz": qz, "qz_m": qz_m, "qz_v": qz_v, "z_expr": z_expr, diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index fde19f3c2f..82f47f4149 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -122,7 +122,8 @@ def __init__( self.cond_cat = conditional_category if conditional_norm and self.n_cat_list[self.cond_cat] == 0: raise ValueError( - "Conditional normalization is not applicable for a categorical variable with only one category." + "Conditional normalization is not applicable for a categorical variable with only " + "one category." ) self.n_cov = n_continuous + sum(self.n_cat_list) From 884983d07f04ce010be717c6ce8d91522ebdf4f2 Mon Sep 17 00:00:00 2001 From: ori-kron-wis Date: Thu, 28 May 2026 12:27:16 +0300 Subject: [PATCH 23/24] reformating this branch of scvi_X --- docs/user_guide/index.md | 3 + docs/user_guide/models/index.md | 1 + docs/user_guide/models/scvix.md | 5 - src/scvi/_constants.py | 2 - src/scvi/external/gimvi/_task.py | 10 +- src/scvi/external/scvix/_model.py | 35 +++-- src/scvi/external/scvix/_module.py | 118 +++++++++++--- src/scvi/external/sysvi/_base_components.py | 6 +- src/scvi/external/sysvi/_module.py | 8 +- src/scvi/model/base/_embedding_mixin.py | 28 +--- src/scvi/module/_constants.py | 3 - src/scvi/module/_multivae.py | 4 +- src/scvi/module/base/_embedding_mixin.py | 68 +------- src/scvi/nn/_base_components.py | 163 +++----------------- src/scvi/nn/_embedding.py | 2 +- src/scvi/train/_trainingplans.py | 92 +++-------- src/scvi/utils/_docstrings.py | 6 - tests/external/scvix/test_scvix.py | 51 ++++++ 18 files changed, 244 insertions(+), 361 deletions(-) diff --git a/docs/user_guide/index.md b/docs/user_guide/index.md index 26d47265a1..733dd740f1 100644 --- a/docs/user_guide/index.md +++ b/docs/user_guide/index.md @@ -21,6 +21,9 @@ scvi-tools is composed of models that can perform one or many analysis tasks. In * - :doc:`/user_guide/models/scvi` - Dimensionality reduction, removal of unwanted variation, integration across replicates, donors, and technologies, differential expression, imputation, normalization of other cell- and sample-level confounding factors - :cite:p:`Lopez18` + * - :doc:`/user_guide/models/scvix` + - Cross-technology integration and zero-shot query mapping for scRNA-seq datasets + - In preparation * - :doc:`/user_guide/models/scanvi` - scVI tasks with cell type transfer from reference, seed labeling - :cite:p:`Xu21` diff --git a/docs/user_guide/models/index.md b/docs/user_guide/models/index.md index e845022bed..31d5fe963d 100644 --- a/docs/user_guide/models/index.md +++ b/docs/user_guide/models/index.md @@ -24,6 +24,7 @@ scanvi scar scbasset scvi +scvix scviva solo stereoscope diff --git a/docs/user_guide/models/scvix.md b/docs/user_guide/models/scvix.md index 48ba92664a..d572b1dbdf 100644 --- a/docs/user_guide/models/scvix.md +++ b/docs/user_guide/models/scvix.md @@ -25,11 +25,6 @@ The key advantages of scVI-X over scVI are: encodings with low-dimensional variational embeddings, reducing parameter count and improving identifiability through a KL regularisation term. -```{topic} Tutorials: - -- {doc}`/tutorials/notebooks/scrna/scvix_tutorial` -``` - ## Preliminaries scVI-X takes as input a scRNA-seq count matrix $X$ with $N$ cells and $G$ diff --git a/src/scvi/_constants.py b/src/scvi/_constants.py index 28caa685f1..ec6e4e914d 100644 --- a/src/scvi/_constants.py +++ b/src/scvi/_constants.py @@ -5,8 +5,6 @@ class _REGISTRY_KEYS_NT(NamedTuple): X_KEY: str = "X" ATAC_X_KEY: str = "atac" BATCH_KEY: str = "batch" - ASSAY_KEY: str = "assay" - ADVERSARIAL_GROUP_KEY: str = "adversarial_group" SAMPLE_KEY: str = "sample" LABELS_KEY: str = "labels" PROTEIN_EXP_KEY: str = "proteins" diff --git a/src/scvi/external/gimvi/_task.py b/src/scvi/external/gimvi/_task.py index d82e4ab9e2..04aa97702d 100644 --- a/src/scvi/external/gimvi/_task.py +++ b/src/scvi/external/gimvi/_task.py @@ -66,10 +66,7 @@ def training_step(self, batch, batch_idx): ] if kappa > 0 and self.adversarial_classifier is not False: fool_loss = self.loss_adversarial_classifier( - torch.cat(zs), - torch.cat(batch_tensor).long(), - torch.cat(batch_tensor).long(), - predict_true_class=False, + torch.cat(zs), torch.cat(batch_tensor).long(), False ) loss += fool_loss * kappa opt1.zero_grad() @@ -97,10 +94,7 @@ def training_step(self, batch, batch_idx): torch.zeros((z.shape[0], 1), device=z.device) + i for i, z in enumerate(zs) ] loss = self.loss_adversarial_classifier( - torch.cat(zs).detach(), - torch.cat(batch_tensor).long(), - torch.cat(batch_tensor).long(), - predict_true_class=True, + torch.cat(zs).detach(), torch.cat(batch_tensor).long(), True ) loss *= kappa opt2.zero_grad() diff --git a/src/scvi/external/scvix/_model.py b/src/scvi/external/scvix/_model.py index 510c545179..ea63248f36 100644 --- a/src/scvi/external/scvix/_model.py +++ b/src/scvi/external/scvix/_model.py @@ -26,10 +26,11 @@ RNASeqMixin, VAEMixin, ) -from scvi.train import AdversarialTrainingPlan, TrainRunner +from scvi.train import TrainRunner from scvi.utils._docstrings import devices_dsp, setup_anndata_dsp from ._module import VAEX +from ._trainingplans import SCVIXTrainingPlan if TYPE_CHECKING: from typing import Literal @@ -37,8 +38,10 @@ import pandas as pd from anndata import AnnData +ASSAY_KEY = "assay" +ADVERSARIAL_GROUP_KEY = "adversarial_group" + logger = logging.getLogger(__name__) -print(2) class SCVIX( @@ -48,7 +51,7 @@ class SCVIX( ArchesMixin, BaseMinifiedModeModelClass, ): - """single-cell Variational Inference :cite:p:`Lopez18`. + """scVI-X: learning technology-invariant cell states. Parameters ---------- @@ -121,7 +124,7 @@ class SCVIX( _LATENT_QZM_KEY = "scvix_latent_qzm" _LATENT_QZV_KEY = "scvix_latent_qzv" _data_splitter_cls = DataSplitter - _training_plan_cls = AdversarialTrainingPlan + _training_plan_cls = SCVIXTrainingPlan _train_runner_cls = TrainRunner def __init__( @@ -157,7 +160,7 @@ def __init__( **kwargs, } self._model_summary_string = ( - "SCVI model with the following parameters: \n" + "SCVI-X model with the following parameters: \n" f"n_hidden: {n_hidden}, n_latent: {n_latent}, n_layers: {n_layers}, " f"dropout_rate: {dropout_rate}, dispersion: {dispersion}, " f"gene_likelihood: {gene_likelihood}, " @@ -169,8 +172,13 @@ def __init__( pseudoinputs_data_indices = np.random.randint( 0, self.summary_stats.n_cells, n_prior_components ) - assert pseudoinputs_data_indices.shape[0] == n_prior_components - assert pseudoinputs_data_indices.ndim == 1 + if pseudoinputs_data_indices.ndim != 1: + raise ValueError("`pseudoinputs_data_indices` must be one-dimensional.") + if pseudoinputs_data_indices.shape[0] != n_prior_components: + raise ValueError( + "`pseudoinputs_data_indices` must contain exactly " + f"{n_prior_components} indices." + ) pseudoinput_data = next( iter( self._make_data_loader( @@ -228,10 +236,10 @@ def train( early_stopping: bool = False, n_epochs_kl_warmup: int | None = 5, adversarial_classifier: bool | None = None, - adversarial_key: str = "assay", + adversarial_key: str = ASSAY_KEY, scale_adversarial_loss: float | str = 5.0, - adversarial_steps: int = 5, - weight_kl_sample: float = 1e-5, + adversarial_steps: int = 1, + weight_kl_sample: float = 0.0, datasplitter_kwargs: dict | None = None, plan_kwargs: dict | None = None, external_indexing: list[np.array] = None, @@ -383,10 +391,9 @@ def get_normalized_expression( adata_ = self._validate_anndata(adata) assay_mappings = ( self.get_anndata_manager(adata_, required=True) - .get_state_registry(REGISTRY_KEYS.ASSAY_KEY) + .get_state_registry(ASSAY_KEY) .categorical_mapping ) - # Hack for transform_assay to be passed in get_normalized_expression. if isinstance(transform_assay, str): if transform_assay not in assay_mappings: raise ValueError(f'"{transform_assay}" is not a valid assay category.') @@ -456,7 +463,7 @@ def setup_anndata( anndata_fields = [ LayerField(REGISTRY_KEYS.X_KEY, layer, is_count_data=True), CategoricalObsField(REGISTRY_KEYS.BATCH_KEY, batch_key), - CategoricalObsField(REGISTRY_KEYS.ASSAY_KEY, assay_key), + CategoricalObsField(ASSAY_KEY, assay_key), CategoricalJointObsField(REGISTRY_KEYS.CAT_COVS_KEY, categorical_covariate_keys), NumericalJointObsField(REGISTRY_KEYS.CONT_COVS_KEY, continuous_covariate_keys), ] @@ -468,7 +475,7 @@ def setup_anndata( ) if adversarial_group_key is not None: anndata_fields.append( - CategoricalObsField(REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, adversarial_group_key) + CategoricalObsField(ADVERSARIAL_GROUP_KEY, adversarial_group_key) ) # register new fields if the adata is minified adata_minify_type = _get_adata_minify_type(adata) diff --git a/src/scvi/external/scvix/_module.py b/src/scvi/external/scvix/_module.py index 94827ddaad..a0001ee4a9 100644 --- a/src/scvi/external/scvix/_module.py +++ b/src/scvi/external/scvix/_module.py @@ -20,8 +20,11 @@ VampPrior, auto_move_data, ) +from scvi.nn import Embedding from scvi.utils import unsupported_if_adata_minified +from ._components import DecoderSCVIX, EncoderX + if TYPE_CHECKING: from collections.abc import Callable from typing import Literal @@ -30,6 +33,11 @@ logger = logging.getLogger(__name__) +ASSAY_KEY = "assay" +ASSAY_INDEX_KEY = "assay_index" +ADVERSARIAL_GROUP_KEY = "adversarial_group" +KL_SAMPLE_KEY = "kl_divergence_sample" + class VAEX(EmbeddingModuleMixin, BaseMinifiedModeModuleClass): """Variational auto-encoder :cite:p:`Lopez18`. @@ -158,8 +166,6 @@ def __init__( pseudoinput_data: dict | None = None, n_prior_components: int | None = 50, ): - from scvi.nn import DecoderSCVI, Encoder - super().__init__() self.dispersion = dispersion @@ -188,6 +194,7 @@ def __init__( ) self.batch_representation = batch_representation + self.batch_representation_encoder = False n_cats_per_cov_ = list([] if n_cats_per_cov is None else n_cats_per_cov) n_continuous = n_continuous_cov @@ -228,7 +235,7 @@ def __init__( if conditional_norm and not encode_assay: conditional_category = 1 _extra_encoder_kwargs = extra_encoder_kwargs or {} - self.z_encoder = Encoder( + self.z_encoder = EncoderX( n_input, n_latent, n_continuous=n_cont_encoder, @@ -247,7 +254,7 @@ def __init__( ) _extra_decoder_kwargs = extra_decoder_kwargs or {} - self.decoder = DecoderSCVI( + self.decoder = DecoderSCVIX( n_latent, n_input, n_cat_list=cat_list, @@ -280,9 +287,8 @@ def __init__( if prior == "gaussian": self.prior = GaussianPrior() elif prior == "vamp": - assert pseudoinput_data is not None, ( - "Pseudoinput data must be specified if using VampPrior" - ) + if pseudoinput_data is None: + raise ValueError("`pseudoinput_data` must be specified if using VampPrior.") pseudoinput_data = self._get_inference_input(pseudoinput_data, full_forward_pass=True) cat_list = [n_batch] + n_cats_per_cov_ + encode_assay_list self.prior = VampPrior( @@ -304,6 +310,75 @@ def __init__( else: raise ValueError("`prior` must be one of 'gaussian', 'vamp', 'mog', 'mog_celltype'.") + @property + def embeddings_dim(self) -> dict[str, int]: + """Dictionary of embedding dimensions before variational doubling.""" + if not hasattr(self, "_embeddings_dim"): + self._embeddings_dim = {} + return self._embeddings_dim + + @property + def embedding_variational(self) -> dict[str, bool]: + """Dictionary of whether each embedding parameterizes a variational distribution.""" + if not hasattr(self, "_embedding_variational"): + self._embedding_variational = {} + return self._embedding_variational + + def init_embedding( + self, + key: str, + num_embeddings: int, + embedding_dim: int = 5, + variational: bool = False, + **kwargs, + ) -> None: + """Initialize a scVI-X embedding without changing the shared embedding mixin API.""" + self.embeddings_dim[key] = embedding_dim + self.embedding_variational[key] = variational + weight_dim = embedding_dim * 2 if variational else embedding_dim + embedding = Embedding(num_embeddings, weight_dim, **kwargs) + torch.nn.init.zeros_(embedding.weight) + self.add_embedding(key, embedding) + + def get_embedding_dim(self, key: str, default_value: int | None = None) -> int: + """Get the non-variational dimension of an embedding.""" + if key not in self.embeddings_dim: + if default_value is not None: + return default_value + raise KeyError(f"Embedding {key} not found.") + return self.embeddings_dim[key] + + def get_embedding_variational(self, key: str, default_value: bool | None = None) -> bool: + """Get whether an embedding is variational.""" + if key not in self.embedding_variational: + if default_value is not None: + return default_value + raise KeyError(f"Embedding {key} not found.") + return self.embedding_variational[key] + + @auto_move_data + def compute_embedding( + self, + key: str, + indices: torch.Tensor, + return_mean: bool = False, + return_dist: bool = False, + ) -> torch.Tensor | distributions.Normal: + """Forward pass for a scVI-X embedding.""" + indices = indices.flatten() if indices.ndim > 1 else indices + embedding = self.get_embedding(key)(indices) + if self.get_embedding_variational(key, default_value=False): + embedding_dim = self.get_embedding_dim(key) + mean = embedding[:, :embedding_dim] + scale = torch.exp(embedding[:, embedding_dim:]) + if return_mean: + return mean + dist = distributions.Normal(mean, scale) + if return_dist: + return dist + return dist.rsample() + return embedding + def _get_inference_input( self, tensors: dict[str, torch.Tensor | None], @@ -324,12 +399,10 @@ def _get_inference_input( return { MODULE_KEYS.X_KEY: tensors[REGISTRY_KEYS.X_KEY], MODULE_KEYS.BATCH_INDEX_KEY: tensors[REGISTRY_KEYS.BATCH_KEY], - MODULE_KEYS.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), + ASSAY_INDEX_KEY: tensors.get(ASSAY_KEY, None), 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.ADVERSARIAL_GROUP_KEY: tensors.get( - REGISTRY_KEYS.ADVERSARIAL_GROUP_KEY, None - ), + ADVERSARIAL_GROUP_KEY: tensors.get(ADVERSARIAL_GROUP_KEY, None), } else: return { @@ -348,7 +421,7 @@ def _get_generative_input( 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.ASSAY_INDEX_KEY: tensors.get(REGISTRY_KEYS.ASSAY_KEY, None), + ASSAY_INDEX_KEY: tensors.get(ASSAY_KEY, None), 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), } @@ -381,6 +454,8 @@ def _regular_inference( batch_rep = self.compute_embedding(REGISTRY_KEYS.BATCH_KEY, batch_index) if cont_covs is not None and self.encode_covariates: cont = torch.cat([cont_covs, batch_rep], dim=-1) + else: + cont = batch_rep else: if self.encode_covariates: cont = cont_covs @@ -401,7 +476,7 @@ def _regular_inference( MODULE_KEYS.Z_KEY: z, MODULE_KEYS.QZ_KEY: qz, MODULE_KEYS.LIBRARY_KEY: library, - MODULE_KEYS.ADVERSARIAL_GROUP_KEY: adversarial_group, + ADVERSARIAL_GROUP_KEY: adversarial_group, } @auto_move_data @@ -417,7 +492,11 @@ def _cached_inference( 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,)) + if n_samples == 1: + untran_z = qz.sample() + else: + qz.sample() + untran_z = qz.sample((n_samples,)) z = self.z_encoder.z_transformation(untran_z) library = torch.log(observed_lib_size) if n_samples > 1: @@ -559,7 +638,7 @@ def loss( else: kl_divergence_sample = torch.zeros_like(kl_divergence_z) - assay_index = tensors.get(REGISTRY_KEYS.ASSAY_KEY, None) + assay_index = tensors.get(ASSAY_KEY, None) if weight_assay_loss > 0.0 and assay_index is not None: assay_loss = self._compute_assay_penalty( inference_outputs[MODULE_KEYS.QZ_KEY].loc, assay_index @@ -598,7 +677,7 @@ def loss( reconstruction_loss=reconst_loss, kl_local={ MODULE_KEYS.KL_Z_KEY: kl_divergence_z, - MODULE_KEYS.KL_SAMPLE_KEY: kl_divergence_sample, + KL_SAMPLE_KEY: kl_divergence_sample, }, ) @@ -664,9 +743,10 @@ def _compute_assay_penalty(self, params, assay): return pair_penalty def mmd(self, params, mask=None): - if mask is not None: - mod_1 = params[mask] - mod_2 = params[~mask] + if mask is None: + raise ValueError("`mask` is required to compute MMD between two groups.") + mod_1 = params[mask] + mod_2 = params[~mask] return rbf_kernel(mod_1, mod_2) diff --git a/src/scvi/external/sysvi/_base_components.py b/src/scvi/external/sysvi/_base_components.py index 375f9dbbb4..81f8d1a859 100644 --- a/src/scvi/external/sysvi/_base_components.py +++ b/src/scvi/external/sysvi/_base_components.py @@ -59,7 +59,7 @@ def __init__( n_input: int, n_output: int, n_cat_list: list[int], - n_continuous: int, + n_cont: int, n_hidden: int = 256, n_layers: int = 3, var_mode: Literal["sample_feature", "feature"] = "feature", @@ -73,7 +73,7 @@ def __init__( self.decoder_y = FCLayers( n_in=n_input, n_cat_list=n_cat_list, - n_continuous=n_continuous, + n_cont=n_cont, n_out=n_hidden, n_hidden=n_hidden, n_layers=n_layers, @@ -116,7 +116,7 @@ def forward( parametrized with the predicted parameters. """ cat_list = [batch_index] + cat_list - q_ = self.decoder_y(x, *cat_list, cont_input=cont) + q_ = self.decoder_y(x, *cat_list, cont=cont) q_m = self.mean_encoder(q_) if q_m.isnan().any() or q_m.isinf().any(): warnings.warn( diff --git a/src/scvi/external/sysvi/_module.py b/src/scvi/external/sysvi/_module.py index 60086ba40e..b06857e439 100644 --- a/src/scvi/external/sysvi/_module.py +++ b/src/scvi/external/sysvi/_module.py @@ -100,13 +100,13 @@ def __init__( self.n_batch = n_batch n_cat_list = [n_batch] - n_continuous = n_continuous_cov + n_cont = n_continuous_cov if n_cats_per_cov is not None: if self.embed_categorical_covariates: for idx, n in enumerate(n_cats_per_cov): covariate_name = f"cov{idx}" self.init_embedding(covariate_name, n, **embedding_kwargs) - n_continuous += self.get_embedding(covariate_name).embedding_dim + n_cont += self.get_embedding(covariate_name).embedding_dim else: n_cat_list.extend(n_cats_per_cov) @@ -114,7 +114,7 @@ def __init__( n_input=n_input, n_output=n_latent, n_cat_list=n_cat_list, - n_continuous=n_continuous, + n_cont=n_cont, n_hidden=n_hidden, n_layers=n_layers, dropout_rate=dropout_rate, @@ -127,7 +127,7 @@ def __init__( n_input=n_latent, n_output=n_input, n_cat_list=n_cat_list, - n_continuous=n_continuous, + n_cont=n_cont, n_hidden=n_hidden, n_layers=n_layers, dropout_rate=dropout_rate, diff --git a/src/scvi/model/base/_embedding_mixin.py b/src/scvi/model/base/_embedding_mixin.py index 03a0626501..7849be92ac 100644 --- a/src/scvi/model/base/_embedding_mixin.py +++ b/src/scvi/model/base/_embedding_mixin.py @@ -13,7 +13,7 @@ class EmbeddingMixin: - """Mixin class for computing covariate embeddings of a model. + """``EXPERIMENTAL`` Mixin class for initializing and using embeddings in a model. Must be used with a module that inherits from :class:`~scvi.module.base.EmbeddingModuleMixin`. @@ -25,34 +25,14 @@ def get_batch_representation( adata: AnnData | None = None, indices: list[int] | None = None, batch_size: int | None = None, - key: str = REGISTRY_KEYS.BATCH_KEY, - return_mean: bool = True, ) -> np.ndarray: - """Get the batch representation for a given set of indices. - - Parameters - ---------- - adata - AnnData object to use. - indices - Indices to get the batch representation for. - batch_size - Minibatch size for computing the batch representation. - key - Setup key to compute the batch representation for. - return_mean - Return the mean of the batch representation. Or sample from it. - """ + """Get the batch representation for a given set of indices.""" if not isinstance(self.module, EmbeddingModuleMixin): raise ValueError("The current `module` must inherit from `EmbeddingModuleMixin`.") - if key not in self.module.embeddings_dim: - raise ValueError(f"Embedding {key} not found. Enable it during model setup.") adata = self._validate_anndata(adata) dataloader = self._make_data_loader(adata=adata, indices=indices, batch_size=batch_size) - tensors = [ - self.module.compute_embedding(key, tensors[key], return_mean=return_mean) - for tensors in dataloader - ] + key = REGISTRY_KEYS.BATCH_KEY + tensors = [self.module.compute_embedding(key, tensors[key]) for tensors in dataloader] return torch.cat(tensors).detach().cpu().numpy() diff --git a/src/scvi/module/_constants.py b/src/scvi/module/_constants.py index 885dc4fe68..2b6e232429 100644 --- a/src/scvi/module/_constants.py +++ b/src/scvi/module/_constants.py @@ -11,8 +11,6 @@ class _MODULE_KEYS(NamedTuple): LIBRARY_KEY: str = "library" QL_KEY: str = "ql" BATCH_INDEX_KEY: str = "batch_index" - ASSAY_INDEX_KEY: str = "assay_index" - ADVERSARIAL_GROUP_KEY: str = "adversarial_group" Y_KEY: str = "y" CONT_COVS_KEY: str = "cont_covs" CAT_COVS_KEY: str = "cat_covs" @@ -24,7 +22,6 @@ class _MODULE_KEYS(NamedTuple): # loss KL_L_KEY: str = "kl_divergence_l" KL_Z_KEY: str = "kl_divergence_z" - KL_SAMPLE_KEY: str = "kl_divergence_sample" MODULE_KEYS = _MODULE_KEYS() diff --git a/src/scvi/module/_multivae.py b/src/scvi/module/_multivae.py index f9ea1c1fdb..58e4efdab9 100644 --- a/src/scvi/module/_multivae.py +++ b/src/scvi/module/_multivae.py @@ -673,14 +673,12 @@ def unsqz(zt, n_s): libsize_acc = unsqz(libsize_acc, n_samples) # sample from the mixed representation - qz = Normal(qz_m, qz_v.sqrt()) - untran_z = qz.rsample() + untran_z = Normal(qz_m, qz_v.sqrt()).rsample() z = self.z_encoder_accessibility.z_transformation(untran_z) outputs = { "x": x, "z": z, - "qz": qz, "qz_m": qz_m, "qz_v": qz_v, "z_expr": z_expr, diff --git a/src/scvi/module/base/_embedding_mixin.py b/src/scvi/module/base/_embedding_mixin.py index 3a58b7d868..1286d27aef 100644 --- a/src/scvi/module/base/_embedding_mixin.py +++ b/src/scvi/module/base/_embedding_mixin.py @@ -1,5 +1,4 @@ import torch -from torch.distributions import Normal from torch.nn import ModuleDict from scvi.module.base._decorators import auto_move_data @@ -7,7 +6,7 @@ class EmbeddingModuleMixin: - """Mixin class for initializing and using embeddings in a module.""" + """``EXPERIMENTAL`` Mixin class for initializing and using embeddings in a module.""" @property def embeddings_dict(self) -> ModuleDict: @@ -16,25 +15,10 @@ def embeddings_dict(self) -> ModuleDict: self._embeddings_dict = ModuleDict() return self._embeddings_dict - @property - def embeddings_dim(self) -> dict: - """Dictionary of embeddings dimensions.""" - if not hasattr(self, "_embeddings_dim"): - self._embeddings_dim = {} - return self._embeddings_dim - - @property - def variational(self) -> dict: - """Dictionary of whether embedding is variational.""" - if not hasattr(self, "_variational"): - self._variational = {} - return self._variational - def add_embedding(self, key: str, embedding: Embedding, overwrite: bool = False) -> None: """Add an embedding to the module.""" if key in self.embeddings_dict and not overwrite: raise KeyError(f"Embedding {key} already exists.") - torch.nn.init.zeros_(embedding.weight) self.embeddings_dict[key] = embedding def remove_embedding(self, key: str) -> None: @@ -43,68 +27,24 @@ def remove_embedding(self, key: str) -> None: raise KeyError(f"Embedding {key} not found.") del self.embeddings_dict[key] - def get_embedding( - self, - key: str, - ) -> Embedding: + def get_embedding(self, key: str) -> Embedding: """Get an embedding from the module.""" if key not in self.embeddings_dict: raise KeyError(f"Embedding {key} not found.") return self.embeddings_dict[key] - def get_embedding_dim(self, key: str, default_value: str | None = None) -> int: - """Get the dimension of an embedding.""" - if key not in self.embeddings_dim: - if default_value is not None: - return default_value - else: - raise KeyError(f"Embedding {key} not found.") - return self.embeddings_dim[key] - - def get_embedding_variational(self, key: str, default_value: str | None = None) -> bool: - """Get whether an embedding is variational.""" - if key not in self.variational: - if default_value is not None: - return default_value - else: - raise KeyError(f"Embedding {key} not found.") - return self.variational[key] - def init_embedding( self, key: str, num_embeddings: int, embedding_dim: int = 5, - variational: bool = False, **kwargs, ) -> None: """Initialize an embedding in the module.""" - self.embeddings_dim[key] = embedding_dim - self.variational[key] = variational - - if variational: - embedding_dim *= 2 self.add_embedding(key, Embedding(num_embeddings, embedding_dim, **kwargs)) @auto_move_data - def compute_embedding( - self, - key: str, - indices: torch.Tensor, - return_mean: bool = False, - return_dist: bool = False, - ) -> torch.Tensor: + def compute_embedding(self, key: str, indices: torch.Tensor) -> torch.Tensor: """Forward pass for an embedding.""" indices = indices.flatten() if indices.ndim > 1 else indices - embedding = self.get_embedding(key)(indices) - if self.get_embedding_variational(key): - embedding_dim = self.get_embedding_dim(key) - if return_mean: - return embedding[:, :embedding_dim] - dist = Normal(embedding[:, :embedding_dim], torch.exp(embedding[:, embedding_dim:])) - if return_dist: - return dist - else: - return dist.rsample() - else: - return embedding + return self.get_embedding(key)(indices) diff --git a/src/scvi/nn/_base_components.py b/src/scvi/nn/_base_components.py index 82f47f4149..fa3e206230 100644 --- a/src/scvi/nn/_base_components.py +++ b/src/scvi/nn/_base_components.py @@ -14,44 +14,6 @@ def _identity(x): return x -class ConditionalBatchNorm2d(nn.Module): - def __init__(self, num_features, num_classes, momentum, eps): - super().__init__() - self.num_features = num_features - self.bn = nn.BatchNorm1d(self.num_features, momentum=momentum, eps=eps, affine=False) - self.embed = nn.Embedding(num_classes, self.num_features * 2) - self.embed.weight.data[:, : self.num_features].normal_( - 1, 0.02 - ) # Initialise scale at N(1, 0.02) - self.embed.weight.data[:, self.num_features :].zero_() # Initialise bias at 0 - - def forward(self, x, y): - out = self.bn(x) - gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) - out = gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) - - return out - - -class ConditionalLayerNorm(nn.Module): - def __init__(self, num_features, num_classes): - super().__init__() - self.num_features = num_features - self.ln = nn.LayerNorm(self.num_features, elementwise_affine=False) - self.embed = nn.Embedding(num_classes, self.num_features * 2) - self.embed.weight.data[:, : self.num_features].normal_( - 1, 0.02 - ) # Initialise scale at N(1, 0.02) - self.embed.weight.data[:, self.num_features :].zero_() # Initialise bias at 0 - - def forward(self, x, y): - out = self.ln(x) - gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) - out = gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) - - return out - - class FCLayers(nn.Module): """A helper class to build fully-connected layers for a neural network. @@ -61,9 +23,6 @@ class FCLayers(nn.Module): The dimensionality of the input n_out The dimensionality of the output - n_continuous - The dimensionality of the continuous covariates - including batch embeddings. n_cat_list A list containing, for each category of interest, the number of categories. Each category will be @@ -94,8 +53,8 @@ def __init__( self, n_in: int, n_out: int, - n_continuous: int = 0, n_cat_list: Iterable[int] = None, + n_cont: int = 0, n_layers: int = 1, n_hidden: int = 128, dropout_rate: float = 0.1, @@ -105,8 +64,6 @@ def __init__( bias: bool = True, inject_covariates: bool = True, activation_fn: nn.Module = nn.ReLU, - conditional_norm: bool = False, - conditional_category: int = 0, ): super().__init__() self.inject_covariates = inject_covariates @@ -117,16 +74,8 @@ def __init__( self.n_cat_list = [n_cat if n_cat > 1 else 0 for n_cat in n_cat_list] else: self.n_cat_list = [] - self.n_continuous = n_continuous - self.cond_cat = conditional_category - if conditional_norm and self.n_cat_list[self.cond_cat] == 0: - raise ValueError( - "Conditional normalization is not applicable for a categorical variable with only " - "one category." - ) - - self.n_cov = n_continuous + sum(self.n_cat_list) + self.n_cov = n_cont + sum(self.n_cat_list) self.fc_layers = nn.Sequential( collections.OrderedDict( @@ -139,18 +88,12 @@ def __init__( n_out, bias=bias, ), - # non-default params come from defaults in Tensorflow implementation - ConditionalBatchNorm2d( - n_out, self.n_cat_list[self.cond_cat], momentum=0.01, eps=0.001 - ) - if conditional_norm and use_batch_norm - else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) + # non-default params come from defaults in the original Tensorflow + # implementation + nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) if use_batch_norm else None, - # non-default params come from defaults in Tensorflow implementation - ConditionalLayerNorm(n_out, self.n_cat_list[self.cond_cat]) - if conditional_norm and use_layer_norm - else nn.LayerNorm(n_out, elementwise_affine=False) + nn.LayerNorm(n_out, elementwise_affine=False) if use_layer_norm else None, activation_fn() if use_activation else None, @@ -196,7 +139,7 @@ def _hook_fn_zero_out(grad): b = layer.bias.register_hook(_hook_fn_zero_out) self.hooks.append(b) - def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | None = None): + def forward(self, x: torch.Tensor, *cat_list: int, cont: torch.Tensor | None = None): """Forward computation on ``x``. Parameters @@ -205,8 +148,8 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No tensor of values with shape ``(n_in,)`` cat_list list of category membership(s) for this sample - cont_input - tensor of continuous covariates with shape ``(n_continuous,)`` + cont + tensor of continuous covariates with shape ``(n_cont,)`` Returns ------- @@ -214,17 +157,11 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No tensor of shape ``(n_out,)`` """ one_hot_cat_list = [] # for generality in this list many idxs useless. - cont_list = [cont_input] if cont_input is not None else [] + cont_list = [cont] if cont is not None else [] cat_list = cat_list or [] if len(self.n_cat_list) > len(cat_list): raise ValueError("nb. categorical args provided doesn't match init. params.") - if ( - self.n_continuous > 0 - and cont_input is not None - and cont_input.shape[-1] != self.n_continuous - ): - raise ValueError("continuous dims provided doesn't match init. params.") for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): if n_cat and cat is None: raise ValueError("cat not provided while n_cat != 0 in init. params.") @@ -238,20 +175,7 @@ def forward(self, x: torch.Tensor, *cat_list: int, cont_input: torch.Tensor | No for i, layers in enumerate(self.fc_layers): for layer in layers: if layer is not None: - if isinstance(layer, ConditionalBatchNorm2d) or isinstance( - layer, ConditionalLayerNorm - ): - if x.dim() == 3: - x = torch.cat( - [ - (layer(x=slice_x, y=cat_list[self.cond_cat])).unsqueeze(0) - for slice_x in x - ], - dim=0, - ) - else: - x = layer(x=x, y=cat_list[self.cond_cat]) - elif isinstance(layer, nn.BatchNorm1d): + if isinstance(layer, nn.BatchNorm1d): if x.dim() == 3: if ( x.device.type == "mps" @@ -291,9 +215,6 @@ class Encoder(nn.Module): The dimensionality of the input (data space) n_output The dimensionality of the output (latent space) - n_continuous - The dimensionality of the continuous covariates - including batch embeddings. n_cat_list A list containing the number of categories for each category of interest. Each category will be @@ -322,8 +243,7 @@ def __init__( self, n_input: int, n_output: int, - n_continuous: int = 0, - n_cat_list: Iterable[int] | None = None, + n_cat_list: Iterable[int] = None, n_layers: int = 1, n_hidden: int = 128, dropout_rate: float = 0.1, @@ -340,7 +260,6 @@ def __init__( self.encoder = FCLayers( n_in=n_input, n_out=n_hidden, - n_continuous=n_continuous, n_cat_list=n_cat_list, n_layers=n_layers, n_hidden=n_hidden, @@ -357,12 +276,7 @@ def __init__( self.z_transformation = _identity self.var_activation = torch.exp if var_activation is None else var_activation - def forward( - self, - x: torch.Tensor, - *cat_list: int, - cont: torch.Tensor | None = None, - ): + def forward(self, x: torch.Tensor, *cat_list: int): r"""The forward computation for a single sample. #. Encodes the data into latent space using the encoder network @@ -376,8 +290,6 @@ def forward( tensor with shape (n_input,) cat_list list of category membership(s) for this sample - cont - optional tensor with shape (n_continuous,) Returns ------- @@ -386,7 +298,7 @@ def forward( """ # Parameters for latent distribution - q = self.encoder(x, *cat_list, cont_input=cont) + q = self.encoder(x, *cat_list) q_m = self.mean_encoder(q) q_v = self.var_activation(self.var_encoder(q)) + self.var_eps dist = Normal(q_m, q_v.sqrt()) @@ -408,8 +320,6 @@ class DecoderSCVI(nn.Module): The dimensionality of the input (latent space) n_output The dimensionality of the output (data space) - n_continuous - The dimensionality of the continuous covariates n_cat_list A list containing the number of categories for each category of interest. Each category will be @@ -418,8 +328,6 @@ class DecoderSCVI(nn.Module): The number of fully-connected hidden layers n_hidden The number of nodes per hidden layer - n_conditions_output - The number of conditions add to the scale and dropout parameters. dropout_rate Dropout rate to apply to each of the hidden layers inject_covariates @@ -438,23 +346,19 @@ def __init__( self, n_input: int, n_output: int, - n_continuous: int = 0, n_cat_list: Iterable[int] = None, n_layers: int = 1, n_hidden: int = 128, - n_conditions_output: int = 0, inject_covariates: bool = True, use_batch_norm: bool = False, use_layer_norm: bool = False, - scale_activation: Literal["softmax", "softplus", "exp"] = "softmax", + scale_activation: Literal["softmax", "softplus"] = "softmax", **kwargs, ): super().__init__() - self.n_conditions_output = n_conditions_output self.px_decoder = FCLayers( n_in=n_input, n_out=n_hidden, - n_continuous=n_continuous, n_cat_list=n_cat_list, n_layers=n_layers, n_hidden=n_hidden, @@ -470,18 +374,16 @@ def __init__( px_scale_activation = nn.Softmax(dim=-1) elif scale_activation == "softplus": px_scale_activation = nn.Softplus() - elif scale_activation == "exp": - px_scale_activation = ExpActivation() - - # scale self.px_scale_decoder = nn.Sequential( - nn.Linear(n_hidden + n_conditions_output, n_output), + nn.Linear(n_hidden, n_output), px_scale_activation, ) - # dispersion: here we only deal with gene-cell dispersion case - self.px_r_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) + + # dispersion: here we only deal with a gene-cell dispersion case + self.px_r_decoder = nn.Linear(n_hidden, n_output) + # dropout - self.px_dropout_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) + self.px_dropout_decoder = nn.Linear(n_hidden, n_output) def forward( self, @@ -489,8 +391,6 @@ def forward( z: torch.Tensor, library: torch.Tensor, *cat_list: int, - cont: torch.Tensor | None = None, - output_condition: torch.Tensor | None = None, ): """The forward computation for a single sample. @@ -513,10 +413,6 @@ def forward( library size cat_list list of category membership(s) for this sample - cont - tensor with shape ``(n_continuous,)`` - output_condition - tensor with shape ``(n_input,)`` used for conditioning the output layer Returns ------- @@ -525,21 +421,12 @@ def forward( """ # The decoder returns values for the parameters of the ZINB distribution - px = self.px_decoder(z, *cat_list, cont_input=cont) - if output_condition is not None and self.n_conditions_output: - one_hot_cat = nn.functional.one_hot( - output_condition.squeeze(-1), self.n_conditions_output - ) - else: - one_hot_cat = torch.zeros(px.size(-2), self.n_conditions_output).to(px.device) - if px.dim() == 3: - one_hot_cat = one_hot_cat.unsqueeze(0).expand(px.size(0), -1, -1) - px_cat = torch.cat([px, one_hot_cat], dim=-1) - px_scale = self.px_scale_decoder(px_cat) - px_dropout = self.px_dropout_decoder(px_cat) + px = self.px_decoder(z, *cat_list) + px_scale = self.px_scale_decoder(px) + px_dropout = self.px_dropout_decoder(px) # Clamp to high value: exp(12) ~ 160000 to avoid nans (computational stability) px_rate = torch.exp(library) * px_scale # torch.clamp( , max=12) - px_r = self.px_r_decoder(px_cat) if dispersion == "gene-cell" else None + px_r = self.px_r_decoder(px) if dispersion == "gene-cell" else None return px_scale, px_r, px_rate, px_dropout diff --git a/src/scvi/nn/_embedding.py b/src/scvi/nn/_embedding.py index 2faf885ab7..58a2740e1c 100644 --- a/src/scvi/nn/_embedding.py +++ b/src/scvi/nn/_embedding.py @@ -27,7 +27,7 @@ def _partial_freeze_hook(grad: torch.Tensor) -> torch.Tensor: class Embedding(nn.Embedding): - """Embedding layer with utility methods for extending.""" + """``EXPERIMENTAL`` Embedding layer with utility methods for extending.""" @classmethod def extend( diff --git a/src/scvi/train/_trainingplans.py b/src/scvi/train/_trainingplans.py index 84002386d1..8373d9d41f 100644 --- a/src/scvi/train/_trainingplans.py +++ b/src/scvi/train/_trainingplans.py @@ -573,10 +573,6 @@ class AdversarialTrainingPlan(TrainingPlan): Minimum learning rate allowed adversarial_classifier Whether to use adversarial classifier in the latent space - adversarial_key - Key in setup args to use for adversarial training. - adversarial_steps - Number of steps to train the adversarial classifier for each training step. scale_adversarial_loss Scaling factor on the adversarial components of the loss. By default, adversarial loss is scaled from 1 to 0 following the opposite of @@ -607,8 +603,6 @@ def __init__( ] = "elbo_validation", lr_min: float = 0, adversarial_classifier: bool | Classifier = False, - adversarial_key: str = "batch", - adversarial_steps: int = 1, scale_adversarial_loss: float | Literal["auto"] = "auto", compile: bool = False, compile_kwargs: dict | None = None, @@ -632,68 +626,43 @@ def __init__( compile_kwargs=compile_kwargs, **loss_kwargs, ) - self.adversarial_key = adversarial_key if adversarial_classifier is True: - self.adversarial_steps = adversarial_steps - self.n_adversarial_name = f"n_{adversarial_key}" - if hasattr(self.module, self.n_adversarial_name): - self.n_output_classifier = getattr(self.module, self.n_adversarial_name) - else: - raise ValueError( - f"Adversarial key {adversarial_key} not found in module setup args." - ) - if self.n_output_classifier == 1: + if self.module.n_batch == 1: warnings.warn( - "Disabling adversarial classifier as there is only one class.", + "Disabling adversarial classifier.", UserWarning, stacklevel=settings.warnings_stacklevel, ) self.adversarial_classifier = False else: + self.n_output_classifier = self.module.n_batch self.adversarial_classifier = Classifier( - n_input=self.module.n_latent + getattr(self.module, "n_adversarial_group", 0), - n_hidden=128, + n_input=self.module.n_latent, + n_hidden=32, n_labels=self.n_output_classifier, n_layers=2, logits=True, - use_batch_norm=False, - use_layer_norm=True, ) - else: self.adversarial_classifier = adversarial_classifier self.scale_adversarial_loss = scale_adversarial_loss self.automatic_optimization = False - def loss_adversarial_classifier( - self, z, adversarial_group, batch_index, predict_true_class=True - ): + def loss_adversarial_classifier(self, z, batch_index, predict_true_class=True): """Loss for adversarial classifier.""" n_classes = self.n_output_classifier - n_adv_group = getattr(self.module, "n_adversarial_group", 0) - if n_adv_group > 0: - adversarial_group_ = torch.nn.functional.one_hot( - adversarial_group, num_classes=n_adv_group - ).float() - z_cls = torch.cat([z, adversarial_group_], dim=1) - else: - z_cls = z - if predict_true_class: # train classifier - z_cls = z_cls.detach() - z = z_cls - cls_logits = self.adversarial_classifier(z) + cls_logits = torch.nn.LogSoftmax(dim=1)(self.adversarial_classifier(z)) if predict_true_class: - cls_target = batch_index.squeeze(-1) - loss = torch.nn.functional.cross_entropy(cls_logits, cls_target) + cls_target = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes) else: - one_hot_batch = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes).float() - cls_target = (1 - one_hot_batch) / (n_classes - 1) - loss = ( - -(cls_target * torch.nn.functional.log_softmax(cls_logits, dim=1)) - .sum(dim=1) - .mean() - ) + one_hot_batch = torch.nn.functional.one_hot(batch_index.squeeze(-1), n_classes) + # place zeroes where the true label is + cls_target = (~one_hot_batch.bool()).float() + cls_target = cls_target / (n_classes - 1) + + l_soft = cls_logits * cls_target + loss = -l_soft.sum(dim=1).mean() return loss @@ -707,7 +676,7 @@ def training_step(self, batch, batch_idx): if self.scale_adversarial_loss == "auto" else self.scale_adversarial_loss ) - batch_tensor = batch[self.adversarial_key].long() + batch_tensor = batch[REGISTRY_KEYS.BATCH_KEY] opts = self.optimizers() if not isinstance(opts, list): @@ -718,16 +687,11 @@ def training_step(self, batch, batch_idx): inference_outputs, _, scvi_loss = self.forward(batch, loss_kwargs=self.loss_kwargs) z = inference_outputs["z"] - adversarial_group = inference_outputs.get("adversarial_group", None) - if adversarial_group is None: - adversarial_group = torch.zeros(z.size(0)).to(z.device).long() - else: - adversarial_group = adversarial_group.squeeze(-1).long() loss = scvi_loss.loss orig_loss = loss # fool classifier if doing adversarial training if kappa > 0 and self.adversarial_classifier is not False: - fool_loss = self.loss_adversarial_classifier(z, adversarial_group, batch_tensor, False) + fool_loss = self.loss_adversarial_classifier(z, batch_tensor, False) loss += fool_loss * kappa self.log("train_loss", loss, on_step=self.on_step, on_epoch=self.on_epoch, prog_bar=True) @@ -744,19 +708,11 @@ def training_step(self, batch, batch_idx): # train adversarial classifier # this condition will not be met unless self.adversarial_classifier is not False if opt2 is not None: - loss = 0.0 - for i in range(self.adversarial_steps): - qz = inference_outputs["qz"] - z = qz.sample() - loss_ = kappa * self.loss_adversarial_classifier( - z, adversarial_group, batch_tensor, True - ) - if i > 1 and (loss - loss_) / loss < 1e-3: - break - loss = loss_ - opt2.zero_grad() - self.manual_backward(loss) - opt2.step() + loss = self.loss_adversarial_classifier(z.detach(), batch_tensor, True) + loss *= kappa + opt2.zero_grad() + self.manual_backward(loss) + opt2.step() # next part is for the usage of scib-metrics autotune with scvi if scvi_loss.extra_metrics is not None and len(scvi_loss.extra_metrics.keys()) > 0: @@ -805,7 +761,9 @@ def configure_optimizers(self): if self.adversarial_classifier is not False: params2 = filter(lambda p: p.requires_grad, self.adversarial_classifier.parameters()) - optimizer2 = torch.optim.Adam(params2, lr=3e-4, eps=1e-4, weight_decay=1e-9) + optimizer2 = torch.optim.Adam( + params2, lr=1e-3, eps=0.01, weight_decay=self.weight_decay + ) config2 = {"optimizer": optimizer2} # pytorch lightning requires this way to return diff --git a/src/scvi/utils/_docstrings.py b/src/scvi/utils/_docstrings.py index 33ad12b0f4..1963cab7cb 100644 --- a/src/scvi/utils/_docstrings.py +++ b/src/scvi/utils/_docstrings.py @@ -124,12 +124,6 @@ integer categories and saved to `adata.obs['_scvi_batch']`. If `None`, assigns the same batch to all the data.""" -param_assay_key = """\ -assay_key - key in `adata.obs` for assay and suspension type information. Categories will automatically be - converted into integer categories and saved to `adata.obs['_scvi_assay']`. If `None`, assigns - the same assay to all the data.""" - param_sample_key = """\ sample_key key in `adata.obs` for sample information. Categories will automatically be converted into diff --git a/tests/external/scvix/test_scvix.py b/tests/external/scvix/test_scvix.py index bd6e47f915..cb1b775eea 100644 --- a/tests/external/scvix/test_scvix.py +++ b/tests/external/scvix/test_scvix.py @@ -2,6 +2,7 @@ import numpy as np import pytest +import torch import scvi from scvi.data import synthetic_iid @@ -103,6 +104,36 @@ def test_scvix_layernorm(): model.get_normalized_expression(n_samples=2) +def test_scvix_batch_representation_encoder_initialized_without_covariates(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch") + model = SCVIX(adata, encode_covariates=False) + + assert model.module.batch_representation_encoder is False + + +def test_scvix_vamp_prior_validates_pseudoinput_shape(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch") + + with pytest.raises(ValueError, match="one-dimensional"): + SCVIX( + adata, + prior="vamp", + pseudoinputs_data_indices=np.array([[0]]), + n_prior_components=1, + ) + + +def test_scvix_mmd_requires_mask(): + adata = synthetic_iid(batch_size=100) + SCVIX.setup_anndata(adata, batch_key="batch", assay_key="batch") + model = SCVIX(adata) + + with pytest.raises(ValueError, match="mask"): + model.module.mmd(torch.randn(4, 2)) + + def test_scvix_scarches_one_hot(save_path): # test transfer_anndata_setup + view adata1 = synthetic_iid() @@ -129,6 +160,16 @@ def test_scvix_scarches_one_hot(save_path): new_var_names = new_var_names_init + adata4.var_names[10:].to_list() adata4.var_names = new_var_names + SCVIX.prepare_query_anndata(adata4, dir_path) + assert np.sum(adata4[:, adata4.var_names[:10]].X) == 0 + np.testing.assert_equal(adata4.var_names[:10].to_numpy(), adata1.var_names[:10].to_numpy()) + SCVIX_query3 = SCVIX.load_query_data(adata4, dir_path) + SCVIX_query3.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + adata5 = SCVIX.prepare_query_anndata(adata4, dir_path, inplace=False) + SCVIX_query4 = SCVIX.load_query_data(adata5, dir_path) + SCVIX_query4.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + def test_scvix_scarches_embedding(save_path): # test transfer_anndata_setup + view @@ -156,6 +197,16 @@ def test_scvix_scarches_embedding(save_path): new_var_names = new_var_names_init + adata4.var_names[10:].to_list() adata4.var_names = new_var_names + SCVIX.prepare_query_anndata(adata4, dir_path) + assert np.sum(adata4[:, adata4.var_names[:10]].X) == 0 + np.testing.assert_equal(adata4.var_names[:10].to_numpy(), adata1.var_names[:10].to_numpy()) + SCVIX_query3 = SCVIX.load_query_data(adata4, dir_path) + SCVIX_query3.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + + adata5 = SCVIX.prepare_query_anndata(adata4, dir_path, inplace=False) + SCVIX_query4 = SCVIX.load_query_data(adata5, dir_path) + SCVIX_query4.train(1, train_size=0.5, plan_kwargs={"weight_decay": 0.0}) + def test_scvix_minified(): adata = synthetic_iid() From 830716855515101ab582a7986672aad8edf2afa6 Mon Sep 17 00:00:00 2001 From: ori-kron-wis Date: Thu, 28 May 2026 13:20:12 +0300 Subject: [PATCH 24/24] adding missing files --- src/scvi/external/scvix/_components.py | 348 ++++++++++++++++++++++ src/scvi/external/scvix/_trainingplans.py | 225 ++++++++++++++ 2 files changed, 573 insertions(+) create mode 100644 src/scvi/external/scvix/_components.py create mode 100644 src/scvi/external/scvix/_trainingplans.py diff --git a/src/scvi/external/scvix/_components.py b/src/scvi/external/scvix/_components.py new file mode 100644 index 0000000000..b232284e59 --- /dev/null +++ b/src/scvi/external/scvix/_components.py @@ -0,0 +1,348 @@ +from __future__ import annotations + +import collections +from typing import TYPE_CHECKING + +import torch +from torch import nn +from torch.distributions import Normal + +if TYPE_CHECKING: + from collections.abc import Callable, Iterable + from typing import Literal + + +def _identity(x): + return x + + +class ConditionalBatchNorm1d(nn.Module): + def __init__(self, num_features: int, num_classes: int, momentum: float, eps: float): + super().__init__() + self.num_features = num_features + self.bn = nn.BatchNorm1d(num_features, momentum=momentum, eps=eps, affine=False) + self.embed = nn.Embedding(num_classes, num_features * 2) + self.embed.weight.data[:, :num_features].normal_(1, 0.02) + self.embed.weight.data[:, num_features:].zero_() + + def forward(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + out = self.bn(x) + gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) + return gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) + + +class ConditionalLayerNorm(nn.Module): + def __init__(self, num_features: int, num_classes: int): + super().__init__() + self.num_features = num_features + self.ln = nn.LayerNorm(num_features, elementwise_affine=False) + self.embed = nn.Embedding(num_classes, num_features * 2) + self.embed.weight.data[:, :num_features].normal_(1, 0.02) + self.embed.weight.data[:, num_features:].zero_() + + def forward(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + out = self.ln(x) + gamma, beta = self.embed(y.long().ravel()).chunk(2, 1) + return gamma.view(-1, self.num_features) * out + beta.view(-1, self.num_features) + + +class FCLayersX(nn.Module): + def __init__( + self, + n_in: int, + n_out: int, + n_continuous: int = 0, + n_cat_list: Iterable[int] | None = None, + n_layers: int = 1, + n_hidden: int = 128, + dropout_rate: float = 0.1, + use_batch_norm: bool = True, + use_layer_norm: bool = False, + use_activation: bool = True, + bias: bool = True, + inject_covariates: bool = True, + activation_fn: Callable[[], nn.Module] = nn.ReLU, + conditional_norm: bool = False, + conditional_category: int = 0, + ): + super().__init__() + self.inject_covariates = inject_covariates + layers_dim = [n_in] + (n_layers - 1) * [n_hidden] + [n_out] + + if n_cat_list is not None: + self.n_cat_list = [n_cat if n_cat > 1 else 0 for n_cat in n_cat_list] + else: + self.n_cat_list = [] + self.n_continuous = n_continuous + self.cond_cat = conditional_category + + if conditional_norm: + invalid_conditional_category = ( + conditional_category >= len(self.n_cat_list) + or self.n_cat_list[conditional_category] == 0 + ) + if invalid_conditional_category: + raise ValueError( + "Conditional normalization requires a categorical covariate with more than " + "one category." + ) + + self.n_cov = n_continuous + sum(self.n_cat_list) + self.fc_layers = nn.Sequential( + collections.OrderedDict( + [ + ( + f"Layer {i}", + nn.Sequential( + nn.Linear( + n_in + self.n_cov * self.inject_into_layer(i), + n_out, + bias=bias, + ), + ConditionalBatchNorm1d( + n_out, + self.n_cat_list[self.cond_cat], + momentum=0.01, + eps=0.001, + ) + if conditional_norm and use_batch_norm + else nn.BatchNorm1d(n_out, momentum=0.01, eps=0.001) + if use_batch_norm + else None, + ConditionalLayerNorm(n_out, self.n_cat_list[self.cond_cat]) + if conditional_norm and use_layer_norm + else nn.LayerNorm(n_out, elementwise_affine=False) + if use_layer_norm + else None, + activation_fn() if use_activation else None, + nn.Dropout(p=dropout_rate) if dropout_rate > 0 else None, + ), + ) + for i, (n_in, n_out) in enumerate( + zip(layers_dim[:-1], layers_dim[1:], strict=True) + ) + ] + ) + ) + + def inject_into_layer(self, layer_num: int) -> bool: + return layer_num == 0 or (layer_num > 0 and self.inject_covariates) + + def set_online_update_hooks(self, hook_first_layer: bool = True): + self.hooks = [] + + def _hook_fn_weight(grad): + categorical_dims = sum(self.n_cat_list) + new_grad = torch.zeros_like(grad) + if categorical_dims > 0: + new_grad[:, -categorical_dims:] = grad[:, -categorical_dims:] + return new_grad + + def _hook_fn_zero_out(grad): + return grad * 0 + + for i, layers in enumerate(self.fc_layers): + for layer in layers: + if i == 0 and not hook_first_layer: + continue + if isinstance(layer, nn.Linear): + if self.inject_into_layer(i): + w = layer.weight.register_hook(_hook_fn_weight) + else: + w = layer.weight.register_hook(_hook_fn_zero_out) + self.hooks.append(w) + b = layer.bias.register_hook(_hook_fn_zero_out) + self.hooks.append(b) + + def forward( + self, + x: torch.Tensor, + *cat_list: int, + cont_input: torch.Tensor | None = None, + ) -> torch.Tensor: + one_hot_cat_list = [] + cont_list = [cont_input] if cont_input is not None else [] + cat_list = cat_list or [] + + if len(self.n_cat_list) > len(cat_list): + raise ValueError("nb. categorical args provided doesn't match init. params.") + if ( + self.n_continuous > 0 + and cont_input is not None + and cont_input.shape[-1] != self.n_continuous + ): + raise ValueError("continuous dims provided doesn't match init. params.") + + for n_cat, cat in zip(self.n_cat_list, cat_list, strict=False): + if n_cat and cat is None: + raise ValueError("cat not provided while n_cat != 0 in init. params.") + if n_cat > 1: + if cat.size(1) != n_cat: + one_hot_cat = nn.functional.one_hot(cat.squeeze(-1), n_cat) + else: + one_hot_cat = cat + one_hot_cat_list.append(one_hot_cat) + cov_list = cont_list + one_hot_cat_list + + for i, layers in enumerate(self.fc_layers): + for layer in layers: + if layer is None: + continue + if isinstance(layer, (ConditionalBatchNorm1d, ConditionalLayerNorm)): + if x.dim() == 3: + x = torch.cat( + [ + layer(slice_x, cat_list[self.cond_cat]).unsqueeze(0) + for slice_x in x + ], + dim=0, + ) + else: + x = layer(x, cat_list[self.cond_cat]) + elif isinstance(layer, nn.BatchNorm1d): + if x.dim() == 3: + if x.device.type == "mps": + x = torch.cat( + [(layer(slice_x.clone())).unsqueeze(0) for slice_x in x], dim=0 + ) + else: + x = torch.cat([layer(slice_x).unsqueeze(0) for slice_x in x], dim=0) + else: + x = layer(x) + else: + if isinstance(layer, nn.Linear) and self.inject_into_layer(i): + if x.dim() == 3: + cov_list_layer = [ + o.unsqueeze(0).expand((x.size(0), o.size(0), o.size(1))) + for o in cov_list + ] + else: + cov_list_layer = cov_list + x = torch.cat((x, *cov_list_layer), dim=-1) + x = layer(x) + return x + + +class EncoderX(nn.Module): + def __init__( + self, + n_input: int, + n_output: int, + n_continuous: int = 0, + n_cat_list: Iterable[int] | None = None, + n_layers: int = 1, + n_hidden: int = 128, + dropout_rate: float = 0.1, + distribution: str = "normal", + var_eps: float = 1e-4, + var_activation: Callable | None = None, + return_dist: bool = False, + **kwargs, + ): + super().__init__() + self.distribution = distribution + self.var_eps = var_eps + self.encoder = FCLayersX( + n_in=n_input, + n_out=n_hidden, + n_continuous=n_continuous, + n_cat_list=n_cat_list, + n_layers=n_layers, + n_hidden=n_hidden, + dropout_rate=dropout_rate, + **kwargs, + ) + self.mean_encoder = nn.Linear(n_hidden, n_output) + self.var_encoder = nn.Linear(n_hidden, n_output) + self.return_dist = return_dist + + self.z_transformation = nn.Softmax(dim=-1) if distribution == "ln" else _identity + self.var_activation = torch.exp if var_activation is None else var_activation + + def forward( + self, + x: torch.Tensor, + *cat_list: int, + cont: torch.Tensor | None = None, + ): + q = self.encoder(x, *cat_list, cont_input=cont) + q_m = self.mean_encoder(q) + q_v = self.var_activation(self.var_encoder(q)) + self.var_eps + dist = Normal(q_m, q_v.sqrt()) + latent = self.z_transformation(dist.rsample()) + if self.return_dist: + return dist, latent + return q_m, q_v, latent + + +class DecoderSCVIX(nn.Module): + def __init__( + self, + n_input: int, + n_output: int, + n_continuous: int = 0, + n_cat_list: Iterable[int] | None = None, + n_layers: int = 1, + n_hidden: int = 128, + n_conditions_output: int = 0, + inject_covariates: bool = True, + use_batch_norm: bool = False, + use_layer_norm: bool = False, + scale_activation: Literal["softmax", "softplus"] = "softmax", + **kwargs, + ): + super().__init__() + self.n_conditions_output = n_conditions_output + self.px_decoder = FCLayersX( + n_in=n_input, + n_out=n_hidden, + n_continuous=n_continuous, + n_cat_list=n_cat_list, + n_layers=n_layers, + n_hidden=n_hidden, + dropout_rate=0, + inject_covariates=inject_covariates, + use_batch_norm=use_batch_norm, + use_layer_norm=use_layer_norm, + **kwargs, + ) + + if scale_activation == "softmax": + px_scale_activation = nn.Softmax(dim=-1) + elif scale_activation == "softplus": + px_scale_activation = nn.Softplus() + else: + raise ValueError("`scale_activation` must be 'softmax' or 'softplus'.") + + self.px_scale_decoder = nn.Sequential( + nn.Linear(n_hidden + n_conditions_output, n_output), + px_scale_activation, + ) + self.px_r_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) + self.px_dropout_decoder = nn.Linear(n_hidden + n_conditions_output, n_output) + + def forward( + self, + dispersion: str, + z: torch.Tensor, + library: torch.Tensor, + *cat_list: int, + cont: torch.Tensor | None = None, + output_condition: torch.Tensor | None = None, + ): + px = self.px_decoder(z, *cat_list, cont_input=cont) + if output_condition is not None and self.n_conditions_output: + one_hot_cat = nn.functional.one_hot( + output_condition.squeeze(-1), self.n_conditions_output + ).float() + else: + one_hot_cat = torch.zeros(px.size(-2), self.n_conditions_output, device=px.device) + if px.dim() == 3: + one_hot_cat = one_hot_cat.unsqueeze(0).expand(px.size(0), -1, -1) + px_cat = torch.cat([px, one_hot_cat], dim=-1) + + px_scale = self.px_scale_decoder(px_cat) + px_dropout = self.px_dropout_decoder(px_cat) + px_rate = torch.exp(library) * px_scale + px_r = self.px_r_decoder(px_cat) if dispersion == "gene-cell" else None + return px_scale, px_r, px_rate, px_dropout diff --git a/src/scvi/external/scvix/_trainingplans.py b/src/scvi/external/scvix/_trainingplans.py new file mode 100644 index 0000000000..590e3df1ce --- /dev/null +++ b/src/scvi/external/scvix/_trainingplans.py @@ -0,0 +1,225 @@ +from __future__ import annotations + +import warnings +from typing import TYPE_CHECKING + +import torch +from torch.optim.lr_scheduler import ReduceLROnPlateau + +from scvi import settings +from scvi.module import Classifier +from scvi.train import AdversarialTrainingPlan + +if TYPE_CHECKING: + from typing import Literal + + from scvi.module.base import BaseModuleClass + from scvi.train._trainingplans import TorchOptimizerCreator + +ASSAY_KEY = "assay" +ADVERSARIAL_GROUP_KEY = "adversarial_group" + + +class SCVIXTrainingPlan(AdversarialTrainingPlan): + """Adversarial training plan with scVI-X-specific classifier inputs.""" + + def __init__( + self, + module: BaseModuleClass, + *, + optimizer: Literal["Adam", "AdamW", "Custom"] = "Adam", + optimizer_creator: TorchOptimizerCreator | None = None, + lr: float = 1e-3, + weight_decay: float = 1e-6, + n_steps_kl_warmup: int = None, + n_epochs_kl_warmup: int = 400, + reduce_lr_on_plateau: bool = False, + lr_factor: float = 0.6, + lr_patience: int = 30, + lr_threshold: float = 0.0, + lr_scheduler_metric: Literal[ + "elbo_validation", "reconstruction_loss_validation", "kl_local_validation" + ] = "elbo_validation", + lr_min: float = 0, + adversarial_classifier: bool | Classifier = False, + adversarial_key: str = ASSAY_KEY, + adversarial_steps: int = 1, + scale_adversarial_loss: float | Literal["auto"] = "auto", + compile: bool = False, + compile_kwargs: dict | None = None, + **loss_kwargs, + ): + super().__init__( + module=module, + optimizer=optimizer, + optimizer_creator=optimizer_creator, + lr=lr, + weight_decay=weight_decay, + n_steps_kl_warmup=n_steps_kl_warmup, + n_epochs_kl_warmup=n_epochs_kl_warmup, + reduce_lr_on_plateau=reduce_lr_on_plateau, + lr_factor=lr_factor, + lr_patience=lr_patience, + lr_threshold=lr_threshold, + lr_scheduler_metric=lr_scheduler_metric, + lr_min=lr_min, + adversarial_classifier=False, + scale_adversarial_loss=scale_adversarial_loss, + compile=compile, + compile_kwargs=compile_kwargs, + **loss_kwargs, + ) + self.adversarial_key = adversarial_key + self.adversarial_steps = adversarial_steps + + if adversarial_classifier is True: + n_adversarial_name = f"n_{adversarial_key}" + if not hasattr(self.module, n_adversarial_name): + raise ValueError( + f"Adversarial key {adversarial_key!r} not found in module setup args." + ) + self.n_output_classifier = getattr(self.module, n_adversarial_name) + if self.n_output_classifier == 1: + warnings.warn( + "Disabling adversarial classifier as there is only one class.", + UserWarning, + stacklevel=settings.warnings_stacklevel, + ) + self.adversarial_classifier = False + else: + self.adversarial_classifier = Classifier( + n_input=self.module.n_latent + getattr(self.module, "n_adversarial_group", 0), + n_hidden=128, + n_labels=self.n_output_classifier, + n_layers=2, + logits=True, + use_batch_norm=False, + use_layer_norm=True, + ) + else: + self.adversarial_classifier = adversarial_classifier + + def loss_adversarial_classifier( + self, + z: torch.Tensor, + adversarial_group: torch.Tensor, + class_index: torch.Tensor, + predict_true_class: bool = True, + ) -> torch.Tensor: + """Loss for the scVI-X adversarial classifier.""" + n_classes = self.n_output_classifier + n_adv_group = getattr(self.module, "n_adversarial_group", 0) + if n_adv_group > 0: + adversarial_group_one_hot = torch.nn.functional.one_hot( + adversarial_group.squeeze(-1).long(), num_classes=n_adv_group + ).float() + z = torch.cat([z, adversarial_group_one_hot], dim=1) + if predict_true_class: + z = z.detach() + + cls_logits = self.adversarial_classifier(z) + if predict_true_class: + return torch.nn.functional.cross_entropy(cls_logits, class_index.squeeze(-1).long()) + + one_hot_class = torch.nn.functional.one_hot(class_index.squeeze(-1), n_classes).float() + cls_target = (1 - one_hot_class) / (n_classes - 1) + return -(cls_target * torch.nn.functional.log_softmax(cls_logits, dim=1)).sum(dim=1).mean() + + def training_step(self, batch, batch_idx): + """Training step for scVI-X adversarial training.""" + if "kl_weight" in self.loss_kwargs: + self.loss_kwargs.update({"kl_weight": self.kl_weight}) + self.log("kl_weight", self.kl_weight, on_step=True, on_epoch=False) + kappa = ( + 1 - self.kl_weight + if self.scale_adversarial_loss == "auto" + else self.scale_adversarial_loss + ) + class_tensor = batch[self.adversarial_key].long() + + opts = self.optimizers() + if not isinstance(opts, list): + opt1 = opts + opt2 = None + else: + opt1, opt2 = opts + + inference_outputs, _, scvi_loss = self.forward(batch, loss_kwargs=self.loss_kwargs) + z = inference_outputs["z"] + adversarial_group = inference_outputs.get(ADVERSARIAL_GROUP_KEY) + if adversarial_group is None: + adversarial_group = torch.zeros(z.size(0), device=z.device, dtype=torch.long) + + loss = scvi_loss.loss + orig_loss = loss + if kappa > 0 and self.adversarial_classifier is not False: + fool_loss = self.loss_adversarial_classifier( + z, + adversarial_group, + class_tensor, + predict_true_class=False, + ) + loss += fool_loss * kappa + + self.log("train_loss", loss, on_step=self.on_step, on_epoch=self.on_epoch, prog_bar=True) + if self.on_step: + self.trainer.logger.log_metrics( + {"train_loss_step": loss}, + step=self.global_step, + ) + self.compute_and_log_metrics(scvi_loss, self.train_metrics, "train") + opt1.zero_grad() + self.manual_backward(loss) + opt1.step() + + if opt2 is not None and kappa > 0: + qz = inference_outputs["qz"] + for _ in range(max(self.adversarial_steps, 0)): + z_classifier = qz.sample() + classifier_loss = kappa * self.loss_adversarial_classifier( + z_classifier, + adversarial_group, + class_tensor, + predict_true_class=True, + ) + opt2.zero_grad() + self.manual_backward(classifier_loss) + opt2.step() + + if scvi_loss.extra_metrics is not None and len(scvi_loss.extra_metrics.keys()) > 0: + self.prepare_scib_autotune(scvi_loss.extra_metrics, "training") + + return orig_loss + + def configure_optimizers(self): + """Configure optimizers for scVI-X adversarial training.""" + params1 = filter(lambda p: p.requires_grad, self.module.parameters()) + optimizer1 = self.get_optimizer_creator()(params1) + config1 = {"optimizer": optimizer1} + if self.reduce_lr_on_plateau: + scheduler1 = ReduceLROnPlateau( + optimizer1, + patience=self.lr_patience, + factor=self.lr_factor, + threshold=self.lr_threshold, + min_lr=self.lr_min, + threshold_mode="abs", + ) + config1.update( + { + "lr_scheduler": { + "scheduler": scheduler1, + "monitor": self.lr_scheduler_metric, + }, + }, + ) + + if self.adversarial_classifier is not False: + params2 = filter(lambda p: p.requires_grad, self.adversarial_classifier.parameters()) + optimizer2 = torch.optim.Adam(params2, lr=3e-4, eps=1e-4, weight_decay=1e-9) + opts = [config1.pop("optimizer"), optimizer2] + if "lr_scheduler" in config1: + return opts, [config1["lr_scheduler"]] + return opts + + return config1