Skip to content

Commit 5794bcb

Browse files
WIP: Update DRVI and align with the latest scvi-tools version (#33)
* Relax dependency versions * Relax scvi dependency and align the code with the newest changes * Allow passing registry and setting adata=None * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Fix deprecated one_hot * Bugfix * Change deprecated csr_matrix .A to .toarray() --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
1 parent 28c9678 commit 5794bcb

10 files changed

Lines changed: 191 additions & 84 deletions

File tree

pyproject.toml

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,19 @@ classifiers = [
3535

3636
# Please make an issue if you need wider range of versions
3737
dependencies = [
38+
"torch>=2.1.0",
39+
"lightning>=2.0",
40+
"scanpy>=1.9.5",
41+
"scikit-learn>=1.5.1",
42+
"scipy>=1.11.3",
43+
"scvi-tools>=1.0.4",
44+
"anndata>=0.10.2",
45+
"numpy>=1.16.1,<2.0.0", # for np.linspace
46+
"pandas>=1.2.0",
47+
# for debug logging (referenced from the issue template)
48+
"session-info",
49+
]
50+
optional-dependencies.restrict = [
3851
"torch>=2.1.0,<2.4",
3952
"lightning>=2.0,<2.1",
4053
"scanpy==1.9.5",
@@ -46,7 +59,7 @@ dependencies = [
4659
"jaxlib<=0.4.20",
4760
## END_TODO
4861
"anndata>=0.10.2,<0.11",
49-
"numpy>=1.16.1", # for np.linspace
62+
"numpy>=1.16.1,<2.0.0", # for np.linspace
5063
"pandas>=1.2.0",
5164
## TODO: update this when this is resolved: https://github.com/boto/botocore/issues/2926
5265
# lightning-cloud depends on boto3 that is currently not compatible with urllib3 so resolution takes forever

src/drvi/nn_modules/prior.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@
55
from torch import nn
66
from torch.distributions import Normal, kl_divergence
77

8+
from drvi.scvi_tools_based.module._constants import MODULE_KEYS
9+
810
# Standard, VaMP, GMM from Karin's CSI repo
911

1012

@@ -198,7 +200,7 @@ def get_params(self) -> tuple[torch.Tensor, torch.Tensor]:
198200
self.encoder.train(False)
199201
if self.input_type == "scfemb":
200202
z = self.encoder({**self.pi_aux_data, **self.pi_tensor_data})
201-
output = z["qz_mean"], z["qz_var"]
203+
output = z[MODULE_KEYS.QZM_KEY], z[MODULE_KEYS.QZV_KEY]
202204
elif self.input_type == "scvi":
203205
if self.preparation_function is None:
204206
raise ValueError("preparation_function must be provided for scvi input type")

src/drvi/scvi_tools_based/model/_drvi.py

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from typing import Any, Literal
44

55
import numpy as np
6+
import scvi
67
from anndata import AnnData
78
from scvi import REGISTRY_KEYS, settings
89
from scvi.data import AnnDataManager
@@ -63,7 +64,8 @@ class DRVI(VAEMixin, DRVIArchesMixin, UnsupervisedTrainingMixin, BaseModelClass,
6364

6465
def __init__(
6566
self,
66-
adata: AnnData | MerlinData,
67+
adata: AnnData | MerlinData | None = None, # TODO: align with all scvi changes: registry, etc.
68+
registry: dict | None = None, # TODO: align with all scvi changes: registry, etc.
6769
n_latent: int = 32,
6870
encoder_dims: Sequence[int] = (128, 128),
6971
decoder_dims: Sequence[int] = (128, 128),
@@ -72,7 +74,10 @@ def __init__(
7274
categorical_covariates: list[str] = (),
7375
**model_kwargs,
7476
) -> None:
75-
super().__init__(adata)
77+
if scvi.__version__ >= "1.3.1":
78+
super().__init__(adata, registry)
79+
else:
80+
super().__init__(adata)
7681

7782
# TODO: Remove later. Currently used to detect autoreload problems sooner.
7883
if isinstance(adata, AnnData):
@@ -87,7 +92,12 @@ def __init__(
8792
)
8893

8994
categorical_covariates_info = FeatureInfoList(categorical_covariates, axis="obs", default_dim=10)
90-
if REGISTRY_KEYS.CAT_COVS_KEY in self.adata_manager.data_registry:
95+
if scvi.__version__ >= "1.3.1" and REGISTRY_KEYS.CAT_COVS_KEY in self.registry["field_registries"]:
96+
cat_cov_stats = self.registry["field_registries"][REGISTRY_KEYS.CAT_COVS_KEY]["state_registry"]
97+
print(cat_cov_stats)
98+
n_cats_per_cov = cat_cov_stats.get("n_cats_per_key", [])
99+
assert tuple(categorical_covariates_info.names) == tuple(cat_cov_stats.get("field_keys", []))
100+
elif scvi.__version__ < "1.3.1" and REGISTRY_KEYS.CAT_COVS_KEY in self.adata_manager.data_registry:
91101
cat_cov_stats = self.adata_manager.get_state_registry(REGISTRY_KEYS.CAT_COVS_KEY)
92102
n_cats_per_cov = cat_cov_stats.n_cats_per_key
93103
assert tuple(categorical_covariates_info.names) == tuple(cat_cov_stats.field_keys)

src/drvi/scvi_tools_based/model/base/_archesmixin.py

Lines changed: 45 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
import logging
22
from collections.abc import Sequence
33

4+
import scvi
45
import torch
56
from anndata import AnnData
7+
from lightning import LightningDataModule
68
from 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
810
from scvi.model._utils import parse_device_args
911
from scvi.model.base import BaseModelClass
1012
from 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:

src/drvi/scvi_tools_based/model/base/_generative_mixin.py

Lines changed: 15 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@
99
from scvi import REGISTRY_KEYS
1010
from torch.nn import functional as F
1111

12+
from drvi.scvi_tools_based.module._constants import MODULE_KEYS
13+
1214
logger = logging.getLogger(__name__)
1315

1416

@@ -93,7 +95,7 @@ def iterate_on_decoded_latent_samples(
9395
>>> import numpy as np
9496
>>> # Define custom step function to extract means
9597
>>> def extract_means(gen_output, store):
96-
... store.append(gen_output["params"]["mean"].detach().cpu())
98+
... store.append(gen_output[MODULE_KEYS.PX_PARAMS_KEY]["mean"].detach().cpu())
9799
>>> # Define aggregation function to concatenate results
98100
>>> def concatenate_results(store):
99101
... return torch.cat(store, dim=0).numpy()
@@ -138,9 +140,9 @@ def iterate_on_decoded_latent_samples(
138140
REGISTRY_KEYS.CAT_COVS_KEY: cat_tensor,
139141
},
140142
inference_outputs={
141-
"z": z_tensor,
142-
"library": lib_tensor,
143-
"gene_likelihood_additional_info": {},
143+
MODULE_KEYS.Z_KEY: z_tensor,
144+
MODULE_KEYS.LIBRARY_KEY: lib_tensor,
145+
MODULE_KEYS.LIKELIHOOD_ADDITIONAL_PARAMS_KEY: {},
144146
},
145147
)
146148
gen_output = self.module.generative(**gen_input)
@@ -209,7 +211,7 @@ def decode_latent_samples(
209211
"""
210212

211213
def step_func(gen_output: dict[str, Any], store: list[Any]) -> None:
212-
store.append(gen_output["params"]["mean"].detach().cpu())
214+
store.append(gen_output[MODULE_KEYS.PX_PARAMS_KEY]["mean"].detach().cpu())
213215

214216
def aggregation_func(store: list[Any]) -> np.ndarray:
215217
return torch.cat(store, dim=0).numpy(force=True)
@@ -375,15 +377,17 @@ def calculate_effect(
375377
inference_outputs: dict[str, Any], generative_outputs: dict[str, Any], losses: Any, store: list[Any]
376378
) -> None:
377379
if self.module.split_aggregation == "logsumexp":
378-
log_mean_params = generative_outputs["original_params"]["mean"] # n_samples x n_splits x n_genes
380+
log_mean_params = generative_outputs[MODULE_KEYS.PX_UNAGGREGATED_PARAMS_KEY][
381+
"mean"
382+
] # n_samples x n_splits x n_genes
379383
log_mean_params = F.pad(
380384
log_mean_params, (0, 0, 0, 1), value=np.log(add_to_counts)
381385
) # n_samples x (n_splits + 1) x n_genes
382386
effect_share = -torch.log(1 - F.softmax(log_mean_params, dim=-2)[:, :-1, :]).sum(
383387
dim=-1
384388
) # n_samples x n_splits
385389
elif self.module.split_aggregation == "sum":
386-
effect_share = torch.abs(generative_outputs["original_params"]["mean"]).sum(
390+
effect_share = torch.abs(generative_outputs[MODULE_KEYS.PX_UNAGGREGATED_PARAMS_KEY]["mean"]).sum(
387391
dim=-1
388392
) # n_samples x n_splits
389393
else:
@@ -474,13 +478,15 @@ def calculate_effect(
474478
inference_outputs: dict[str, Any], generative_outputs: dict[str, Any], losses: Any, store: list[Any]
475479
) -> None:
476480
if self.module.split_aggregation == "logsumexp":
477-
log_mean_params = generative_outputs["original_params"]["mean"] # n_samples x n_splits x n_genes
481+
log_mean_params = generative_outputs[MODULE_KEYS.PX_UNAGGREGATED_PARAMS_KEY][
482+
"mean"
483+
] # n_samples x n_splits x n_genes
478484
log_mean_params = F.pad(
479485
log_mean_params, (0, 0, 0, 1), value=np.log(add_to_counts)
480486
) # n_samples x (n_splits + 1) x n_genes
481487
effect_share = -torch.log(1 - F.softmax(log_mean_params, dim=-2)[:, :-1, :])
482488
elif self.module.split_aggregation == "sum":
483-
effect_share = torch.abs(generative_outputs["original_params"]["mean"])
489+
effect_share = torch.abs(generative_outputs[MODULE_KEYS.PX_UNAGGREGATED_PARAMS_KEY]["mean"])
484490
else:
485491
raise NotImplementedError("Only logsumexp and sum aggregations are supported for now.")
486492
effect_share = effect_share.amax(dim=0).detach().cpu().numpy(force=True)
Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,48 @@
1+
# For backward compatibility
2+
try:
3+
from scvi.module._constants import _MODULE_KEYS as _SCVI_MODULE_KEYS
4+
except ImportError:
5+
from typing import NamedTuple
6+
7+
class _NEW_SCVI_MODULE_KEYS(NamedTuple):
8+
X_KEY: str = "x"
9+
# inference
10+
Z_KEY: str = "z"
11+
QZ_KEY: str = "qz"
12+
QZM_KEY: str = "qzm"
13+
QZV_KEY: str = "qzv"
14+
LIBRARY_KEY: str = "library"
15+
QL_KEY: str = "ql"
16+
BATCH_INDEX_KEY: str = "batch_index"
17+
Y_KEY: str = "y"
18+
CONT_COVS_KEY: str = "cont_covs"
19+
CAT_COVS_KEY: str = "cat_covs"
20+
SIZE_FACTOR_KEY: str = "size_factor"
21+
# generative
22+
PX_KEY: str = "px"
23+
PL_KEY: str = "pl"
24+
PZ_KEY: str = "pz"
25+
# loss
26+
KL_L_KEY: str = "kl_divergence_l"
27+
KL_Z_KEY: str = "kl_divergence_z"
28+
29+
class _SCVI_MODULE_KEYS(_NEW_SCVI_MODULE_KEYS):
30+
QZM_KEY: str = "qz_m"
31+
QZV_KEY: str = "qz_v"
32+
33+
34+
class _DRVI_MODULE_KEYS(_SCVI_MODULE_KEYS):
35+
# generative
36+
PX_PARAMS_KEY = "px_params"
37+
PX_UNAGGREGATED_PARAMS_KEY = "px_unaggregated_params"
38+
# Extra
39+
LIKELIHOOD_ADDITIONAL_PARAMS_KEY: str = "gene_likelihood_additional_info"
40+
X_MASK_KEY: str = "x_mask"
41+
# Tensor IO structure
42+
CONT_COVS_TENSOR_KEY: str = "cont_full_tensor"
43+
CAT_COVS_TENSOR_KEY: str = "cat_full_tensor"
44+
# Loss
45+
MSE_LOSS_KEY: str = "mse"
46+
47+
48+
MODULE_KEYS = _DRVI_MODULE_KEYS()

0 commit comments

Comments
 (0)