diffbio 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- diffbio/__init__.py +39 -0
- diffbio/configs.py +75 -0
- diffbio/constants.py +204 -0
- diffbio/core/__init__.py +127 -0
- diffbio/core/base_operators.py +612 -0
- diffbio/core/data_types.py +260 -0
- diffbio/core/gnn_components.py +629 -0
- diffbio/core/graph_utils.py +149 -0
- diffbio/core/neural_components.py +270 -0
- diffbio/core/optimal_transport.py +133 -0
- diffbio/core/soft_ops/__init__.py +216 -0
- diffbio/core/soft_ops/_projections_permutahedron.py +1864 -0
- diffbio/core/soft_ops/_projections_simplex.py +240 -0
- diffbio/core/soft_ops/_projections_transport.py +508 -0
- diffbio/core/soft_ops/_sorting_network.py +204 -0
- diffbio/core/soft_ops/_types.py +15 -0
- diffbio/core/soft_ops/_utils.py +342 -0
- diffbio/core/soft_ops/autograd_safe.py +120 -0
- diffbio/core/soft_ops/comparison.py +235 -0
- diffbio/core/soft_ops/elementwise.py +309 -0
- diffbio/core/soft_ops/logical.py +146 -0
- diffbio/core/soft_ops/quantile.py +376 -0
- diffbio/core/soft_ops/selection.py +236 -0
- diffbio/core/soft_ops/sorting.py +926 -0
- diffbio/core/soft_ops/straight_through.py +261 -0
- diffbio/core/uncertainty.py +279 -0
- diffbio/evaluation/__init__.py +42 -0
- diffbio/evaluation/adapters.py +409 -0
- diffbio/evaluation/graders.py +223 -0
- diffbio/evaluation/problem.py +157 -0
- diffbio/evaluation/runner.py +277 -0
- diffbio/losses/__init__.py +59 -0
- diffbio/losses/alignment_losses.py +222 -0
- diffbio/losses/biological_regularization.py +288 -0
- diffbio/losses/metric_losses.py +139 -0
- diffbio/losses/singlecell_losses.py +387 -0
- diffbio/losses/statistical_losses.py +345 -0
- diffbio/operators/__init__.py +60 -0
- diffbio/operators/_count_vae.py +197 -0
- diffbio/operators/_loss_balancing.py +65 -0
- diffbio/operators/_masked_gene_transformer.py +118 -0
- diffbio/operators/_transformer_validation.py +50 -0
- diffbio/operators/alignment/__init__.py +51 -0
- diffbio/operators/alignment/profile_hmm.py +350 -0
- diffbio/operators/alignment/scoring.py +127 -0
- diffbio/operators/alignment/smith_waterman.py +261 -0
- diffbio/operators/alignment/soft_msa.py +419 -0
- diffbio/operators/assembly/__init__.py +27 -0
- diffbio/operators/assembly/gnn_assembly.py +252 -0
- diffbio/operators/assembly/metagenomic_binning.py +296 -0
- diffbio/operators/crispr/__init__.py +17 -0
- diffbio/operators/crispr/guide_scoring.py +269 -0
- diffbio/operators/drug_discovery/__init__.py +133 -0
- diffbio/operators/drug_discovery/_graph_utils.py +142 -0
- diffbio/operators/drug_discovery/admet_predictor.py +285 -0
- diffbio/operators/drug_discovery/attentive_fp.py +411 -0
- diffbio/operators/drug_discovery/dti.py +261 -0
- diffbio/operators/drug_discovery/fingerprint.py +490 -0
- diffbio/operators/drug_discovery/maccs_keys.py +267 -0
- diffbio/operators/drug_discovery/message_passing.py +200 -0
- diffbio/operators/drug_discovery/primitives.py +242 -0
- diffbio/operators/drug_discovery/property_predictor.py +163 -0
- diffbio/operators/drug_discovery/similarity.py +193 -0
- diffbio/operators/epigenomics/__init__.py +35 -0
- diffbio/operators/epigenomics/chromatin_state.py +491 -0
- diffbio/operators/epigenomics/contextual.py +288 -0
- diffbio/operators/epigenomics/fno_peak_calling.py +153 -0
- diffbio/operators/epigenomics/peak_calling.py +555 -0
- diffbio/operators/foundation_models/__init__.py +119 -0
- diffbio/operators/foundation_models/adapters.py +114 -0
- diffbio/operators/foundation_models/contracts.py +245 -0
- diffbio/operators/foundation_models/embedding_probe.py +83 -0
- diffbio/operators/foundation_models/experimental.py +128 -0
- diffbio/operators/foundation_models/foundation_model.py +332 -0
- diffbio/operators/foundation_models/frozen.py +59 -0
- diffbio/operators/foundation_models/precomputed.py +270 -0
- diffbio/operators/foundation_models/transformer_encoder.py +564 -0
- diffbio/operators/mapping/__init__.py +17 -0
- diffbio/operators/mapping/neural_mapper.py +493 -0
- diffbio/operators/metabolomics/__init__.py +39 -0
- diffbio/operators/metabolomics/spectral_similarity.py +315 -0
- diffbio/operators/molecular_dynamics/__init__.py +51 -0
- diffbio/operators/molecular_dynamics/force_field.py +265 -0
- diffbio/operators/molecular_dynamics/integrator.py +304 -0
- diffbio/operators/molecular_dynamics/primitives.py +115 -0
- diffbio/operators/multiomics/__init__.py +38 -0
- diffbio/operators/multiomics/hic_contact.py +377 -0
- diffbio/operators/multiomics/multiomics_vae.py +325 -0
- diffbio/operators/multiomics/spatial_deconvolution.py +316 -0
- diffbio/operators/multiomics/spatial_gene_detection.py +493 -0
- diffbio/operators/normalization/__init__.py +42 -0
- diffbio/operators/normalization/embedding.py +222 -0
- diffbio/operators/normalization/phate.py +400 -0
- diffbio/operators/normalization/umap.py +261 -0
- diffbio/operators/normalization/vae_normalizer.py +258 -0
- diffbio/operators/population/__init__.py +17 -0
- diffbio/operators/population/ancestry_estimation.py +274 -0
- diffbio/operators/preprocessing/__init__.py +76 -0
- diffbio/operators/preprocessing/adapter_removal.py +311 -0
- diffbio/operators/preprocessing/duplicate_filter.py +317 -0
- diffbio/operators/preprocessing/error_correction.py +287 -0
- diffbio/operators/protein/__init__.py +31 -0
- diffbio/operators/protein/secondary_structure.py +509 -0
- diffbio/operators/quality_filter.py +128 -0
- diffbio/operators/rna_structure/__init__.py +35 -0
- diffbio/operators/rna_structure/rna_folding.py +509 -0
- diffbio/operators/rnaseq/__init__.py +23 -0
- diffbio/operators/rnaseq/motif_discovery.py +251 -0
- diffbio/operators/rnaseq/splicing_psi.py +216 -0
- diffbio/operators/singlecell/__init__.py +193 -0
- diffbio/operators/singlecell/ambient_removal.py +333 -0
- diffbio/operators/singlecell/archetypes.py +191 -0
- diffbio/operators/singlecell/batch_correction.py +288 -0
- diffbio/operators/singlecell/cell_annotation.py +519 -0
- diffbio/operators/singlecell/communication.py +704 -0
- diffbio/operators/singlecell/differential_distribution.py +243 -0
- diffbio/operators/singlecell/doublet_detection.py +657 -0
- diffbio/operators/singlecell/downsampling.py +166 -0
- diffbio/operators/singlecell/enhanced_batch_correction.py +519 -0
- diffbio/operators/singlecell/grn_inference.py +336 -0
- diffbio/operators/singlecell/imputation.py +429 -0
- diffbio/operators/singlecell/knockdown_filter.py +176 -0
- diffbio/operators/singlecell/ot_trajectory.py +277 -0
- diffbio/operators/singlecell/simulation.py +444 -0
- diffbio/operators/singlecell/sindy_grn.py +247 -0
- diffbio/operators/singlecell/soft_clustering.py +211 -0
- diffbio/operators/singlecell/spatial_domains.py +677 -0
- diffbio/operators/singlecell/switch_de.py +184 -0
- diffbio/operators/singlecell/trajectory.py +447 -0
- diffbio/operators/singlecell/velocity.py +361 -0
- diffbio/operators/statistical/__init__.py +35 -0
- diffbio/operators/statistical/em_quantification.py +260 -0
- diffbio/operators/statistical/hmm.py +234 -0
- diffbio/operators/statistical/nb_glm.py +272 -0
- diffbio/operators/variant/__init__.py +64 -0
- diffbio/operators/variant/classifier.py +333 -0
- diffbio/operators/variant/cnn_classifier.py +255 -0
- diffbio/operators/variant/cnv_segmentation.py +678 -0
- diffbio/operators/variant/deepvariant_pileup.py +426 -0
- diffbio/operators/variant/pileup.py +240 -0
- diffbio/operators/variant/quality_recalibration.py +274 -0
- diffbio/pipelines/__init__.py +65 -0
- diffbio/pipelines/differential_expression.py +279 -0
- diffbio/pipelines/enhanced_variant_calling.py +326 -0
- diffbio/pipelines/perturbation.py +407 -0
- diffbio/pipelines/preprocessing.py +267 -0
- diffbio/pipelines/single_cell.py +366 -0
- diffbio/pipelines/variant_calling.py +490 -0
- diffbio/samplers/__init__.py +9 -0
- diffbio/samplers/perturbation_sampler.py +142 -0
- diffbio/sequences/__init__.py +34 -0
- diffbio/sequences/dna.py +239 -0
- diffbio/sources/__init__.py +149 -0
- diffbio/sources/_anndata_shared.py +89 -0
- diffbio/sources/_batch_iteration.py +37 -0
- diffbio/sources/_benchmark_source.py +152 -0
- diffbio/sources/_indexed_batch_source.py +38 -0
- diffbio/sources/_utils.py +45 -0
- diffbio/sources/anndata_interop.py +387 -0
- diffbio/sources/anndata_source.py +361 -0
- diffbio/sources/archive_ii.py +174 -0
- diffbio/sources/balifam.py +207 -0
- diffbio/sources/bam.py +265 -0
- diffbio/sources/bengrn_ground_truth.py +306 -0
- diffbio/sources/contextual_epigenomics.py +242 -0
- diffbio/sources/dti.py +359 -0
- diffbio/sources/embeddings.py +203 -0
- diffbio/sources/encode_peaks.py +223 -0
- diffbio/sources/fasta.py +226 -0
- diffbio/sources/immune_human.py +172 -0
- diffbio/sources/indexed_embeddings.py +128 -0
- diffbio/sources/indexed_view.py +191 -0
- diffbio/sources/molnet.py +493 -0
- diffbio/sources/multiomics.py +279 -0
- diffbio/sources/pancreas.py +108 -0
- diffbio/sources/perturbation/__init__.py +69 -0
- diffbio/sources/perturbation/_types.py +51 -0
- diffbio/sources/perturbation/_utils.py +125 -0
- diffbio/sources/perturbation/concat_source.py +115 -0
- diffbio/sources/perturbation/control_mapping.py +215 -0
- diffbio/sources/perturbation/experiment_config.py +261 -0
- diffbio/sources/perturbation/h5_metadata_cache.py +218 -0
- diffbio/sources/perturbation/output_space.py +52 -0
- diffbio/sources/perturbation/perturbation_source.py +513 -0
- diffbio/sources/seqfish.py +145 -0
- diffbio/sources/sequence_foundation.py +68 -0
- diffbio/sources/singlecell_foundation.py +68 -0
- diffbio/splitters/__init__.py +63 -0
- diffbio/splitters/base.py +251 -0
- diffbio/splitters/molecular.py +330 -0
- diffbio/splitters/perturbation.py +199 -0
- diffbio/splitters/random.py +217 -0
- diffbio/splitters/sequence.py +201 -0
- diffbio/utils/__init__.py +55 -0
- diffbio/utils/dependency_runtime.py +115 -0
- diffbio/utils/nn_utils.py +157 -0
- diffbio/utils/quality.py +45 -0
- diffbio/utils/training.py +585 -0
- diffbio-0.1.0.dist-info/METADATA +480 -0
- diffbio-0.1.0.dist-info/RECORD +202 -0
- diffbio-0.1.0.dist-info/WHEEL +4 -0
- diffbio-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,519 @@
|
|
|
1
|
+
"""Cell type annotation operator for single-cell analysis.
|
|
2
|
+
|
|
3
|
+
This module provides a differentiable cell type annotator supporting three
|
|
4
|
+
annotation strategies inspired by popular tools:
|
|
5
|
+
|
|
6
|
+
- **celltypist**: Logistic-regression classifier on a VAE latent space.
|
|
7
|
+
- **cellassign**: Marker-gene likelihood model with learnable rate parameters.
|
|
8
|
+
- **scanvi**: Semi-supervised VAE with type-conditioned latent prior.
|
|
9
|
+
For each cell type y, learns prior parameters mu_y and logvar_y so that
|
|
10
|
+
``KL(q(z|x) || p(z|y))`` encourages different types to occupy distinct
|
|
11
|
+
latent regions. For unlabelled cells the KL is marginalised over
|
|
12
|
+
predicted type probabilities.
|
|
13
|
+
|
|
14
|
+
All three modes are end-to-end differentiable and JIT-compatible, enabling
|
|
15
|
+
gradient-based optimisation of annotation models within a Datarax pipeline.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
import logging
|
|
19
|
+
from dataclasses import dataclass, field
|
|
20
|
+
from typing import Any, Literal
|
|
21
|
+
|
|
22
|
+
import jax
|
|
23
|
+
import jax.numpy as jnp
|
|
24
|
+
from artifex.generative_models.core.losses.divergence import gaussian_kl_divergence
|
|
25
|
+
from datarax.core.config import OperatorConfig
|
|
26
|
+
from flax import nnx
|
|
27
|
+
from jaxtyping import Array, Float, Int, PyTree
|
|
28
|
+
|
|
29
|
+
from diffbio.constants import EPSILON
|
|
30
|
+
|
|
31
|
+
from diffbio.core.base_operators import EncoderDecoderOperator
|
|
32
|
+
from diffbio.operators._count_vae import CountReconstructionMixin, CountVAEBackboneMixin
|
|
33
|
+
from diffbio.utils.nn_utils import get_rng_key
|
|
34
|
+
|
|
35
|
+
logger = logging.getLogger(__name__)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True)
|
|
39
|
+
class CellAnnotatorConfig(OperatorConfig):
|
|
40
|
+
"""Configuration for cell type annotation.
|
|
41
|
+
|
|
42
|
+
Attributes:
|
|
43
|
+
annotation_mode: Annotation strategy to use.
|
|
44
|
+
n_cell_types: Number of cell types to classify.
|
|
45
|
+
n_genes: Number of input genes.
|
|
46
|
+
latent_dim: Latent-space dimensionality for VAE encoder.
|
|
47
|
+
hidden_dims: Hidden layer sizes for encoder and decoder.
|
|
48
|
+
marker_matrix_shape: Shape (n_types, n_genes) for cellassign mode.
|
|
49
|
+
gene_likelihood: Reconstruction likelihood for scanvi mode.
|
|
50
|
+
``"poisson"`` for standard Poisson NLL (default),
|
|
51
|
+
``"zinb"`` for Zero-Inflated Negative Binomial.
|
|
52
|
+
"""
|
|
53
|
+
|
|
54
|
+
annotation_mode: Literal["scanvi", "cellassign", "celltypist"] = "celltypist"
|
|
55
|
+
n_cell_types: int = 10
|
|
56
|
+
n_genes: int = 2000
|
|
57
|
+
latent_dim: int = 10
|
|
58
|
+
hidden_dims: list[int] = field(default_factory=lambda: [128, 64])
|
|
59
|
+
marker_matrix_shape: tuple[int, int] | None = None
|
|
60
|
+
gene_likelihood: Literal["poisson", "zinb"] = "poisson"
|
|
61
|
+
|
|
62
|
+
def __post_init__(self) -> None:
|
|
63
|
+
"""Set stochastic defaults and validate."""
|
|
64
|
+
object.__setattr__(self, "stochastic", True)
|
|
65
|
+
if self.stream_name is None:
|
|
66
|
+
object.__setattr__(self, "stream_name", "sample")
|
|
67
|
+
super().__post_init__()
|
|
68
|
+
|
|
69
|
+
if self.n_cell_types <= 0:
|
|
70
|
+
raise ValueError(f"n_cell_types must be positive, got {self.n_cell_types}")
|
|
71
|
+
if self.n_genes <= 0:
|
|
72
|
+
raise ValueError(f"n_genes must be positive, got {self.n_genes}")
|
|
73
|
+
if self.latent_dim <= 0:
|
|
74
|
+
raise ValueError(f"latent_dim must be positive, got {self.latent_dim}")
|
|
75
|
+
if any(dim <= 0 for dim in self.hidden_dims):
|
|
76
|
+
raise ValueError(
|
|
77
|
+
f"hidden_dims must contain only positive values, got {self.hidden_dims}"
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
expected_marker_shape = (self.n_cell_types, self.n_genes)
|
|
81
|
+
if self.annotation_mode == "cellassign":
|
|
82
|
+
if self.marker_matrix_shape is None:
|
|
83
|
+
raise ValueError(
|
|
84
|
+
"marker_matrix_shape must be provided for annotation_mode='cellassign'"
|
|
85
|
+
)
|
|
86
|
+
if self.marker_matrix_shape != expected_marker_shape:
|
|
87
|
+
raise ValueError(
|
|
88
|
+
"marker_matrix_shape must match "
|
|
89
|
+
f"(n_cell_types, n_genes)={expected_marker_shape}, "
|
|
90
|
+
f"got {self.marker_matrix_shape}"
|
|
91
|
+
)
|
|
92
|
+
elif self.marker_matrix_shape is not None:
|
|
93
|
+
raise ValueError(
|
|
94
|
+
"marker_matrix_shape is only supported for annotation_mode='cellassign'"
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
if self.annotation_mode != "scanvi" and self.gene_likelihood != "poisson":
|
|
98
|
+
raise ValueError("gene_likelihood is only configurable for annotation_mode='scanvi'")
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
class DifferentiableCellAnnotator(
|
|
102
|
+
CountReconstructionMixin,
|
|
103
|
+
CountVAEBackboneMixin,
|
|
104
|
+
EncoderDecoderOperator,
|
|
105
|
+
):
|
|
106
|
+
"""Differentiable cell type annotator with three annotation modes.
|
|
107
|
+
|
|
108
|
+
Modes
|
|
109
|
+
-----
|
|
110
|
+
**celltypist** (logistic regression on latent):
|
|
111
|
+
Encode counts to a VAE latent, apply a linear classifier head, softmax.
|
|
112
|
+
|
|
113
|
+
**cellassign** (marker-gene likelihood):
|
|
114
|
+
Given a binary marker matrix *M*, compute per-type Poisson
|
|
115
|
+
log-likelihoods with learnable rate parameters, then softmax.
|
|
116
|
+
|
|
117
|
+
**scanvi** (semi-supervised VAE with type-conditioned prior):
|
|
118
|
+
VAE encoder + classifier head with learnable per-type Gaussian priors
|
|
119
|
+
in latent space. The KL divergence uses ``p(z|y) = N(mu_y, sigma_y)``
|
|
120
|
+
instead of the standard ``N(0, I)``, and is marginalised over predicted
|
|
121
|
+
type probabilities for unlabelled cells.
|
|
122
|
+
|
|
123
|
+
All modes additionally produce a latent representation via a shared
|
|
124
|
+
VAE encoder.
|
|
125
|
+
|
|
126
|
+
Inherits from EncoderDecoderOperator to get:
|
|
127
|
+
|
|
128
|
+
- reparameterize() for the VAE sampling step
|
|
129
|
+
- kl_divergence() for KL from standard normal
|
|
130
|
+
- elbo_loss() for combining reconstruction and KL losses
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
config: CellAnnotatorConfig with model parameters.
|
|
134
|
+
rngs: Flax NNX random number generators.
|
|
135
|
+
name: Optional operator name.
|
|
136
|
+
|
|
137
|
+
Example:
|
|
138
|
+
```python
|
|
139
|
+
config = CellAnnotatorConfig(
|
|
140
|
+
annotation_mode="celltypist",
|
|
141
|
+
n_cell_types=10,
|
|
142
|
+
n_genes=2000,
|
|
143
|
+
stochastic=True,
|
|
144
|
+
stream_name="sample",
|
|
145
|
+
)
|
|
146
|
+
annotator = DifferentiableCellAnnotator(config, rngs=nnx.Rngs(42))
|
|
147
|
+
data = {"counts": counts}
|
|
148
|
+
result, state, meta = annotator.apply(data, {}, None)
|
|
149
|
+
```
|
|
150
|
+
"""
|
|
151
|
+
|
|
152
|
+
def __init__(
|
|
153
|
+
self,
|
|
154
|
+
config: CellAnnotatorConfig,
|
|
155
|
+
*,
|
|
156
|
+
rngs: nnx.Rngs | None = None,
|
|
157
|
+
name: str | None = None,
|
|
158
|
+
) -> None:
|
|
159
|
+
"""Initialise the cell type annotator.
|
|
160
|
+
|
|
161
|
+
Args:
|
|
162
|
+
config: Annotator configuration.
|
|
163
|
+
rngs: Random number generators for initialisation and sampling.
|
|
164
|
+
name: Optional operator name.
|
|
165
|
+
"""
|
|
166
|
+
super().__init__(config, rngs=rngs, name=name)
|
|
167
|
+
|
|
168
|
+
rngs = self._init_count_vae_operator(config=config, rngs=rngs)
|
|
169
|
+
|
|
170
|
+
# --- mode-specific heads ---
|
|
171
|
+
if config.annotation_mode in ("celltypist", "scanvi"):
|
|
172
|
+
self.classifier_head = nnx.Linear(
|
|
173
|
+
in_features=config.latent_dim,
|
|
174
|
+
out_features=config.n_cell_types,
|
|
175
|
+
rngs=rngs,
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
if config.annotation_mode == "scanvi":
|
|
179
|
+
# Type-conditioned prior parameters: each cell type y has its own
|
|
180
|
+
# Gaussian prior N(mu_y, diag(exp(logvar_y))) in latent space.
|
|
181
|
+
params_key = get_rng_key(rngs, "params", fallback_seed=7)
|
|
182
|
+
self.prior_means = nnx.Param(
|
|
183
|
+
jax.random.normal(params_key, (config.n_cell_types, config.latent_dim)) * 0.01
|
|
184
|
+
)
|
|
185
|
+
self.prior_logvars = nnx.Param(jnp.zeros((config.n_cell_types, config.latent_dim)))
|
|
186
|
+
|
|
187
|
+
if config.annotation_mode == "scanvi" and config.gene_likelihood == "zinb":
|
|
188
|
+
# ZINB decoder heads: log-dispersion and dropout logit
|
|
189
|
+
# Decoder reverses hidden_dims, so the final hidden dim is the first
|
|
190
|
+
last_hidden = config.hidden_dims[0] if config.hidden_dims else config.latent_dim
|
|
191
|
+
self.fc_log_theta = nnx.Linear(
|
|
192
|
+
in_features=last_hidden,
|
|
193
|
+
out_features=config.n_genes,
|
|
194
|
+
rngs=rngs,
|
|
195
|
+
)
|
|
196
|
+
self.fc_pi_logit = nnx.Linear(
|
|
197
|
+
in_features=last_hidden,
|
|
198
|
+
out_features=config.n_genes,
|
|
199
|
+
rngs=rngs,
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
if config.annotation_mode == "cellassign":
|
|
203
|
+
# Learnable log-rate parameters: mu_type_g (in log space).
|
|
204
|
+
# Initialise to log(5) so Poisson rates start at ~5;
|
|
205
|
+
# the x*log(mu) term is then sensitive to count magnitude,
|
|
206
|
+
# letting the marker matrix drive type discrimination.
|
|
207
|
+
self.log_mu = nnx.Param(
|
|
208
|
+
jnp.full(
|
|
209
|
+
(config.n_cell_types, config.n_genes),
|
|
210
|
+
jnp.log(5.0),
|
|
211
|
+
)
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
def decode(
|
|
215
|
+
self,
|
|
216
|
+
z: Float[Array, "batch latent_dim"],
|
|
217
|
+
) -> dict[str, Float[Array, "batch n_genes"]]:
|
|
218
|
+
"""Decode latent vectors to gene expression parameters.
|
|
219
|
+
|
|
220
|
+
Args:
|
|
221
|
+
z: Latent representations, shape ``(n, latent_dim)``.
|
|
222
|
+
|
|
223
|
+
Returns:
|
|
224
|
+
Dictionary with ``"log_rate"`` (always present) and optionally
|
|
225
|
+
``"log_theta"`` and ``"pi_logit"`` when ZINB likelihood is active.
|
|
226
|
+
"""
|
|
227
|
+
x = self.decode_hidden(z)
|
|
228
|
+
|
|
229
|
+
result: dict[str, Float[Array, "batch n_genes"]] = {
|
|
230
|
+
"log_rate": self.fc_output(x),
|
|
231
|
+
}
|
|
232
|
+
|
|
233
|
+
if self.config.gene_likelihood == "zinb":
|
|
234
|
+
result["log_theta"] = self.fc_log_theta(x)
|
|
235
|
+
result["pi_logit"] = self.fc_pi_logit(x)
|
|
236
|
+
|
|
237
|
+
return result
|
|
238
|
+
|
|
239
|
+
# ------------------------------------------------------------------
|
|
240
|
+
# Per-mode annotation logic
|
|
241
|
+
# ------------------------------------------------------------------
|
|
242
|
+
|
|
243
|
+
def _annotate_celltypist(
|
|
244
|
+
self,
|
|
245
|
+
z: Float[Array, "batch latent_dim"],
|
|
246
|
+
) -> Float[Array, "batch n_cell_types"]:
|
|
247
|
+
"""Celltypist: logistic classifier on latent.
|
|
248
|
+
|
|
249
|
+
Args:
|
|
250
|
+
z: Latent representations.
|
|
251
|
+
|
|
252
|
+
Returns:
|
|
253
|
+
Cell type probabilities, shape ``(n, n_cell_types)``.
|
|
254
|
+
"""
|
|
255
|
+
logits = self.classifier_head(z)
|
|
256
|
+
return jax.nn.softmax(logits, axis=-1)
|
|
257
|
+
|
|
258
|
+
def _annotate_cellassign(
|
|
259
|
+
self,
|
|
260
|
+
counts: Float[Array, "batch n_genes"],
|
|
261
|
+
marker_matrix: Float[Array, "n_types n_genes"],
|
|
262
|
+
) -> Float[Array, "batch n_cell_types"]:
|
|
263
|
+
"""Cellassign: marker-gene Poisson likelihood.
|
|
264
|
+
|
|
265
|
+
For each cell type, compute masked Poisson log-likelihood using only
|
|
266
|
+
the marker genes for that type.
|
|
267
|
+
|
|
268
|
+
Args:
|
|
269
|
+
counts: Gene expression counts, shape ``(n, n_genes)``.
|
|
270
|
+
marker_matrix: Binary marker matrix, shape ``(n_types, n_genes)``.
|
|
271
|
+
|
|
272
|
+
Returns:
|
|
273
|
+
Cell type probabilities, shape ``(n, n_cell_types)``.
|
|
274
|
+
"""
|
|
275
|
+
# Learnable rates in positive space
|
|
276
|
+
mu = jnp.exp(self.log_mu[...]) # (n_types, n_genes)
|
|
277
|
+
|
|
278
|
+
# Poisson log-likelihood per gene per type:
|
|
279
|
+
# log P(x_g | type) = x_g * log(mu_type_g) - mu_type_g - log(x_g!)
|
|
280
|
+
# We mask by M so only marker genes contribute.
|
|
281
|
+
# shape: (1, n_genes) vs (n_types, n_genes)
|
|
282
|
+
log_mu = jnp.log(mu + EPSILON) # (n_types, n_genes)
|
|
283
|
+
|
|
284
|
+
# counts: (n, g), log_mu: (t, g), marker: (t, g)
|
|
285
|
+
# Expand counts: (n, 1, g)
|
|
286
|
+
counts_expanded = counts[:, None, :] # (n, 1, g)
|
|
287
|
+
|
|
288
|
+
# Per-gene log-likelihood (ignoring log(x!) which cancels in softmax)
|
|
289
|
+
# ll_g = x_g * log(mu) - mu, masked by marker
|
|
290
|
+
log_lik_per_gene = counts_expanded * log_mu[None, :, :] - mu[None, :, :]
|
|
291
|
+
# Mask by marker matrix
|
|
292
|
+
masked_log_lik = log_lik_per_gene * marker_matrix[None, :, :] # (n, t, g)
|
|
293
|
+
# Sum over genes to get per-type log-likelihood
|
|
294
|
+
log_lik = jnp.sum(masked_log_lik, axis=-1) # (n, t)
|
|
295
|
+
|
|
296
|
+
return jax.nn.softmax(log_lik, axis=-1)
|
|
297
|
+
|
|
298
|
+
def _type_conditioned_kl(
|
|
299
|
+
self,
|
|
300
|
+
mean: Float[Array, "batch latent_dim"],
|
|
301
|
+
logvar: Float[Array, "batch latent_dim"],
|
|
302
|
+
type_probs: Float[Array, "batch n_cell_types"],
|
|
303
|
+
) -> Float[Array, ""]:
|
|
304
|
+
"""KL divergence with type-conditioned prior, marginalised over types.
|
|
305
|
+
|
|
306
|
+
For each cell type y with prior ``N(mu_y, diag(exp(logvar_y)))``,
|
|
307
|
+
compute the analytic KL from ``q(z|x) = N(mean, diag(exp(logvar)))``
|
|
308
|
+
then marginalise:
|
|
309
|
+
|
|
310
|
+
KL = sum_y p(y) * KL( q(z|x) || p(z|y) )
|
|
311
|
+
|
|
312
|
+
Args:
|
|
313
|
+
mean: Encoder mean, shape ``(n, latent_dim)``.
|
|
314
|
+
logvar: Encoder log-variance, shape ``(n, latent_dim)``.
|
|
315
|
+
type_probs: Cell-type probabilities, shape ``(n, n_cell_types)``.
|
|
316
|
+
|
|
317
|
+
Returns:
|
|
318
|
+
Scalar marginalised KL divergence (summed over batch).
|
|
319
|
+
"""
|
|
320
|
+
prior_mu = self.prior_means[...] # (n_types, latent_dim)
|
|
321
|
+
prior_lv = self.prior_logvars[...] # (n_types, latent_dim)
|
|
322
|
+
|
|
323
|
+
# Expand for broadcasting:
|
|
324
|
+
# mean/logvar: (n, 1, d), prior: (1, t, d)
|
|
325
|
+
mean_e = mean[:, None, :] # (n, 1, d)
|
|
326
|
+
logvar_e = logvar[:, None, :] # (n, 1, d)
|
|
327
|
+
prior_mu_e = prior_mu[None, :, :] # (1, t, d)
|
|
328
|
+
prior_lv_e = prior_lv[None, :, :] # (1, t, d)
|
|
329
|
+
|
|
330
|
+
# Analytic KL between two diagonal Gaussians per dimension:
|
|
331
|
+
# KL(N(m1,s1)||N(m2,s2))
|
|
332
|
+
# = 0.5*(lv2-lv1 + (exp(lv1)+(m1-m2)^2)/exp(lv2) - 1)
|
|
333
|
+
kl_per_dim = 0.5 * (
|
|
334
|
+
prior_lv_e
|
|
335
|
+
- logvar_e
|
|
336
|
+
+ (jnp.exp(logvar_e) + (mean_e - prior_mu_e) ** 2) / (jnp.exp(prior_lv_e) + EPSILON)
|
|
337
|
+
- 1.0
|
|
338
|
+
) # (n, t, d)
|
|
339
|
+
|
|
340
|
+
kl_per_type = jnp.sum(kl_per_dim, axis=-1) # (n, t)
|
|
341
|
+
|
|
342
|
+
# Marginalise over types: sum_y p(y) * KL_y
|
|
343
|
+
kl_marginal = jnp.sum(type_probs * kl_per_type, axis=-1) # (n,)
|
|
344
|
+
|
|
345
|
+
return jnp.sum(kl_marginal)
|
|
346
|
+
|
|
347
|
+
def _annotate_scanvi(
|
|
348
|
+
self,
|
|
349
|
+
z: Float[Array, "batch latent_dim"],
|
|
350
|
+
mean: Float[Array, "batch latent_dim"],
|
|
351
|
+
logvar: Float[Array, "batch latent_dim"],
|
|
352
|
+
counts: Float[Array, "batch n_genes"],
|
|
353
|
+
known_labels: Int[Array, "n_labeled"] | None,
|
|
354
|
+
label_indices: Int[Array, "n_labeled"] | None,
|
|
355
|
+
) -> Float[Array, "batch n_cell_types"]:
|
|
356
|
+
"""Scanvi: semi-supervised VAE with type-conditioned prior.
|
|
357
|
+
|
|
358
|
+
The classifier head produces type probabilities from latent z.
|
|
359
|
+
For labelled cells the known labels are used directly as one-hot
|
|
360
|
+
type probabilities. For unlabelled cells the predicted
|
|
361
|
+
probabilities are returned as-is.
|
|
362
|
+
|
|
363
|
+
Args:
|
|
364
|
+
z: Latent representations.
|
|
365
|
+
mean: Encoder mean.
|
|
366
|
+
logvar: Encoder log-variance.
|
|
367
|
+
counts: Original counts (for reconstruction context).
|
|
368
|
+
known_labels: Known integer labels for a subset of cells.
|
|
369
|
+
label_indices: Indices into the batch for the labelled cells.
|
|
370
|
+
|
|
371
|
+
Returns:
|
|
372
|
+
Cell type probabilities, shape ``(n, n_cell_types)``.
|
|
373
|
+
"""
|
|
374
|
+
logits = self.classifier_head(z) # (n, n_types)
|
|
375
|
+
probs = jax.nn.softmax(logits, axis=-1)
|
|
376
|
+
|
|
377
|
+
if known_labels is None or label_indices is None:
|
|
378
|
+
return probs
|
|
379
|
+
|
|
380
|
+
# For labelled cells, set probabilities to one-hot of the known type.
|
|
381
|
+
one_hot_targets = jax.nn.one_hot(known_labels, self.config.n_cell_types)
|
|
382
|
+
probs = probs.at[label_indices].set(one_hot_targets)
|
|
383
|
+
|
|
384
|
+
return probs
|
|
385
|
+
|
|
386
|
+
# ------------------------------------------------------------------
|
|
387
|
+
# Training loss (scanvi ELBO)
|
|
388
|
+
# ------------------------------------------------------------------
|
|
389
|
+
|
|
390
|
+
def compute_elbo_loss(
|
|
391
|
+
self,
|
|
392
|
+
counts: Float[Array, "batch n_genes"],
|
|
393
|
+
known_labels: Int[Array, "n_labeled"] | None = None,
|
|
394
|
+
label_indices: Int[Array, "n_labeled"] | None = None,
|
|
395
|
+
beta: float = 1.0,
|
|
396
|
+
) -> Float[Array, ""]:
|
|
397
|
+
"""Compute the negative ELBO for training.
|
|
398
|
+
|
|
399
|
+
For **scanvi** mode the KL term uses a type-conditioned prior:
|
|
400
|
+
each cell type y has its own Gaussian prior ``N(mu_y, sigma_y)``
|
|
401
|
+
and the KL is marginalised over predicted/known type probabilities.
|
|
402
|
+
|
|
403
|
+
For other modes (celltypist, cellassign) the standard
|
|
404
|
+
``KL(q(z|x) || N(0,I))`` is used via artifex.
|
|
405
|
+
|
|
406
|
+
When ``gene_likelihood="zinb"`` (scanvi only), the reconstruction
|
|
407
|
+
loss uses the ZINB negative log-likelihood instead of Poisson NLL.
|
|
408
|
+
|
|
409
|
+
Loss = reconstruction + beta * KL + cross_entropy_on_labeled.
|
|
410
|
+
|
|
411
|
+
Args:
|
|
412
|
+
counts: Gene expression counts ``(n, n_genes)``.
|
|
413
|
+
known_labels: Integer labels for labelled subset (scanvi).
|
|
414
|
+
label_indices: Batch indices of labelled cells (scanvi).
|
|
415
|
+
beta: KL weight (default 1.0, >1 for beta-VAE).
|
|
416
|
+
|
|
417
|
+
Returns:
|
|
418
|
+
Scalar negative ELBO loss.
|
|
419
|
+
"""
|
|
420
|
+
mean, logvar = self.encode(counts)
|
|
421
|
+
z = self.reparameterize(mean, logvar)
|
|
422
|
+
|
|
423
|
+
# Reconstruction loss (Poisson or ZINB depending on config)
|
|
424
|
+
decode_output = self.decode(z)
|
|
425
|
+
recon_loss = self.reconstruction_loss(counts, decode_output)
|
|
426
|
+
|
|
427
|
+
# KL divergence -- type-conditioned for scanvi, standard otherwise
|
|
428
|
+
if self.config.annotation_mode == "scanvi":
|
|
429
|
+
logits = self.classifier_head(z)
|
|
430
|
+
type_probs = jax.nn.softmax(logits, axis=-1)
|
|
431
|
+
|
|
432
|
+
# For labelled cells, override predicted probs with known labels
|
|
433
|
+
if known_labels is not None and label_indices is not None:
|
|
434
|
+
one_hot_known = jax.nn.one_hot(known_labels, self.config.n_cell_types)
|
|
435
|
+
type_probs = type_probs.at[label_indices].set(one_hot_known)
|
|
436
|
+
|
|
437
|
+
kl = self._type_conditioned_kl(mean, logvar, type_probs)
|
|
438
|
+
else:
|
|
439
|
+
kl = gaussian_kl_divergence(mean, logvar, reduction="sum")
|
|
440
|
+
|
|
441
|
+
loss = recon_loss + beta * kl
|
|
442
|
+
|
|
443
|
+
# Cross-entropy on labelled cells (scanvi / celltypist)
|
|
444
|
+
if (
|
|
445
|
+
self.config.annotation_mode in ("scanvi", "celltypist")
|
|
446
|
+
and known_labels is not None
|
|
447
|
+
and label_indices is not None
|
|
448
|
+
):
|
|
449
|
+
logits = self.classifier_head(z)
|
|
450
|
+
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
451
|
+
labeled_log_probs = log_probs[label_indices]
|
|
452
|
+
one_hot = jax.nn.one_hot(known_labels, self.config.n_cell_types)
|
|
453
|
+
ce = -jnp.sum(one_hot * labeled_log_probs)
|
|
454
|
+
loss = loss + ce
|
|
455
|
+
|
|
456
|
+
return loss
|
|
457
|
+
|
|
458
|
+
# ------------------------------------------------------------------
|
|
459
|
+
# apply()
|
|
460
|
+
# ------------------------------------------------------------------
|
|
461
|
+
|
|
462
|
+
def apply(
|
|
463
|
+
self,
|
|
464
|
+
data: PyTree,
|
|
465
|
+
state: PyTree,
|
|
466
|
+
metadata: dict[str, Any] | None,
|
|
467
|
+
random_params: Any = None,
|
|
468
|
+
stats: dict[str, Any] | None = None,
|
|
469
|
+
) -> tuple[PyTree, PyTree, dict[str, Any] | None]:
|
|
470
|
+
"""Annotate cells with type probabilities.
|
|
471
|
+
|
|
472
|
+
Args:
|
|
473
|
+
data: Dictionary containing:
|
|
474
|
+
- ``"counts"``: Gene expression counts ``(n, n_genes)``
|
|
475
|
+
- (cellassign) ``"marker_matrix"``: Binary ``(n_types, n_genes)``
|
|
476
|
+
- (scanvi) ``"known_labels"``: Integer labels ``(n_labeled,)``
|
|
477
|
+
- (scanvi) ``"label_indices"``: Batch indices ``(n_labeled,)``
|
|
478
|
+
state: Element state (passed through unchanged).
|
|
479
|
+
metadata: Element metadata (passed through unchanged).
|
|
480
|
+
random_params: Not used.
|
|
481
|
+
stats: Not used.
|
|
482
|
+
|
|
483
|
+
Returns:
|
|
484
|
+
Tuple of (transformed_data, state, metadata) where
|
|
485
|
+
transformed_data adds:
|
|
486
|
+
- ``"cell_type_probabilities"``: ``(n, n_cell_types)``
|
|
487
|
+
- ``"cell_type_labels"``: ``(n,)`` argmax labels
|
|
488
|
+
- ``"latent"``: ``(n, latent_dim)``
|
|
489
|
+
"""
|
|
490
|
+
counts = data["counts"]
|
|
491
|
+
|
|
492
|
+
# Shared VAE encoding
|
|
493
|
+
mean, logvar = self.encode(counts)
|
|
494
|
+
z = self.reparameterize(mean, logvar)
|
|
495
|
+
|
|
496
|
+
# Mode dispatch
|
|
497
|
+
mode = self.config.annotation_mode
|
|
498
|
+
if mode == "celltypist":
|
|
499
|
+
probs = self._annotate_celltypist(z)
|
|
500
|
+
elif mode == "cellassign":
|
|
501
|
+
marker_matrix = data["marker_matrix"]
|
|
502
|
+
probs = self._annotate_cellassign(counts, marker_matrix)
|
|
503
|
+
elif mode == "scanvi":
|
|
504
|
+
known_labels = data.get("known_labels")
|
|
505
|
+
label_indices = data.get("label_indices")
|
|
506
|
+
probs = self._annotate_scanvi(z, mean, logvar, counts, known_labels, label_indices)
|
|
507
|
+
else:
|
|
508
|
+
raise ValueError(f"Unknown annotation_mode: {mode}")
|
|
509
|
+
|
|
510
|
+
labels = jnp.argmax(probs, axis=-1)
|
|
511
|
+
|
|
512
|
+
transformed_data = {
|
|
513
|
+
**data,
|
|
514
|
+
"cell_type_probabilities": probs,
|
|
515
|
+
"cell_type_labels": labels,
|
|
516
|
+
"latent": z,
|
|
517
|
+
}
|
|
518
|
+
|
|
519
|
+
return transformed_data, state, metadata
|