11import logging
22from collections .abc import Sequence
33
4+ import scvi
45import torch
56from anndata import AnnData
7+ from lightning import LightningDataModule
68from scvi import REGISTRY_KEYS
7- from scvi .data ._constants import _MODEL_NAME_KEY , _SETUP_ARGS_KEY
9+ from scvi .data ._constants import _MODEL_NAME_KEY , _SETUP_ARGS_KEY , _SETUP_METHOD_NAME
810from scvi .model ._utils import parse_device_args
911from scvi .model .base import BaseModelClass
1012from scvi .model .base ._archesmixin import ArchesMixin , _get_loaded_data , _initialize_model , _validate_var_names
@@ -21,8 +23,9 @@ class DRVIArchesMixin(ArchesMixin):
2123 @classmethod
2224 def load_query_data (
2325 cls ,
24- adata : AnnData ,
25- reference_model : str | BaseModelClass ,
26+ adata : AnnData = None ,
27+ reference_model : str | BaseModelClass = None ,
28+ registry : dict = None ,
2629 inplace_subset_query_vars : bool = False ,
2730 accelerator : str = "auto" ,
2831 device : int | str = "auto" ,
@@ -35,6 +38,7 @@ def load_query_data(
3538 reset_decoder : bool = False ,
3639 freeze_batchnorm_encoder : bool = True ,
3740 freeze_batchnorm_decoder : bool = False ,
41+ datamodule : LightningDataModule | None = None ,
3842 ):
3943 """Online update of a reference model with scArches algorithm :cite:p:`Lotfollahi21`.
4044
@@ -67,36 +71,56 @@ def load_query_data(
6771 freeze_batchnorm_decoder
6872 Whether to freeze decoder batchnorms' weight and bias during transfer
6973 """
74+ if reference_model is None :
75+ raise ValueError ("Please provide a reference model as string or loaded model." )
76+ if adata is None and registry is None :
77+ raise ValueError ("Please provide either an AnnData or a registry dictionary." )
78+
7079 _ , _ , device = parse_device_args (
7180 accelerator = accelerator ,
7281 devices = device ,
7382 return_device = "torch" ,
7483 validate_single_device = True ,
7584 )
7685
77- attr_dict , var_names , load_state_dict = _get_loaded_data (reference_model , device = device )
86+ # We limit to [:3] as from scvi version 1.1.5 additional output (pyro_param_store) is returned
87+ attr_dict , var_names , load_state_dict = _get_loaded_data (reference_model , device = device )[:3 ]
7888
79- if inplace_subset_query_vars :
80- logger .debug ("Subsetting query vars to reference vars." )
81- adata ._inplace_subset_var (var_names )
82- _validate_var_names (adata , var_names )
89+ if adata :
90+ if inplace_subset_query_vars :
91+ logger .debug ("Subsetting query vars to reference vars." )
92+ adata ._inplace_subset_var (var_names )
93+ _validate_var_names (adata , var_names )
8394
84- registry = attr_dict .pop ("registry_" )
85- if _MODEL_NAME_KEY in registry and registry [_MODEL_NAME_KEY ] != cls .__name__ :
86- raise ValueError ("It appears you are loading a model from a different class." )
95+ registry = attr_dict .pop ("registry_" )
96+ if _MODEL_NAME_KEY in registry and registry [_MODEL_NAME_KEY ] != cls .__name__ :
97+ raise ValueError ("It appears you are loading a model from a different class." )
8798
88- if _SETUP_ARGS_KEY not in registry :
89- raise ValueError ("Saved model does not contain original setup inputs. Cannot load the original setup." )
99+ if _SETUP_ARGS_KEY not in registry :
100+ raise ValueError ("Saved model does not contain original setup inputs. Cannot load the original setup." )
90101
91- cls .setup_anndata (
92- adata ,
93- source_registry = registry ,
94- extend_categories = True ,
95- allow_missing_labels = True ,
96- ** registry [_SETUP_ARGS_KEY ],
97- )
102+ if registry [_SETUP_METHOD_NAME ] != "setup_datamodule" :
103+ setup_method = getattr (cls , registry [_SETUP_METHOD_NAME ])
104+ setup_method (
105+ adata ,
106+ source_registry = registry ,
107+ extend_categories = True ,
108+ allow_missing_labels = True ,
109+ ** registry [_SETUP_ARGS_KEY ],
110+ )
98111
99- model = _initialize_model (cls , adata , attr_dict )
112+ cls .setup_anndata (
113+ adata ,
114+ source_registry = registry ,
115+ extend_categories = True ,
116+ allow_missing_labels = True ,
117+ ** registry [_SETUP_ARGS_KEY ],
118+ )
119+
120+ if scvi .__version__ >= "1.3.1" :
121+ model = _initialize_model (cls , adata , registry , attr_dict , datamodule )
122+ else :
123+ model = _initialize_model (cls , adata , attr_dict )
100124 adata_manager = model .get_anndata_manager (adata , required = True )
101125
102126 if REGISTRY_KEYS .CAT_COVS_KEY in adata_manager .data_registry :
0 commit comments