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.
Files changed (202) hide show
  1. diffbio/__init__.py +39 -0
  2. diffbio/configs.py +75 -0
  3. diffbio/constants.py +204 -0
  4. diffbio/core/__init__.py +127 -0
  5. diffbio/core/base_operators.py +612 -0
  6. diffbio/core/data_types.py +260 -0
  7. diffbio/core/gnn_components.py +629 -0
  8. diffbio/core/graph_utils.py +149 -0
  9. diffbio/core/neural_components.py +270 -0
  10. diffbio/core/optimal_transport.py +133 -0
  11. diffbio/core/soft_ops/__init__.py +216 -0
  12. diffbio/core/soft_ops/_projections_permutahedron.py +1864 -0
  13. diffbio/core/soft_ops/_projections_simplex.py +240 -0
  14. diffbio/core/soft_ops/_projections_transport.py +508 -0
  15. diffbio/core/soft_ops/_sorting_network.py +204 -0
  16. diffbio/core/soft_ops/_types.py +15 -0
  17. diffbio/core/soft_ops/_utils.py +342 -0
  18. diffbio/core/soft_ops/autograd_safe.py +120 -0
  19. diffbio/core/soft_ops/comparison.py +235 -0
  20. diffbio/core/soft_ops/elementwise.py +309 -0
  21. diffbio/core/soft_ops/logical.py +146 -0
  22. diffbio/core/soft_ops/quantile.py +376 -0
  23. diffbio/core/soft_ops/selection.py +236 -0
  24. diffbio/core/soft_ops/sorting.py +926 -0
  25. diffbio/core/soft_ops/straight_through.py +261 -0
  26. diffbio/core/uncertainty.py +279 -0
  27. diffbio/evaluation/__init__.py +42 -0
  28. diffbio/evaluation/adapters.py +409 -0
  29. diffbio/evaluation/graders.py +223 -0
  30. diffbio/evaluation/problem.py +157 -0
  31. diffbio/evaluation/runner.py +277 -0
  32. diffbio/losses/__init__.py +59 -0
  33. diffbio/losses/alignment_losses.py +222 -0
  34. diffbio/losses/biological_regularization.py +288 -0
  35. diffbio/losses/metric_losses.py +139 -0
  36. diffbio/losses/singlecell_losses.py +387 -0
  37. diffbio/losses/statistical_losses.py +345 -0
  38. diffbio/operators/__init__.py +60 -0
  39. diffbio/operators/_count_vae.py +197 -0
  40. diffbio/operators/_loss_balancing.py +65 -0
  41. diffbio/operators/_masked_gene_transformer.py +118 -0
  42. diffbio/operators/_transformer_validation.py +50 -0
  43. diffbio/operators/alignment/__init__.py +51 -0
  44. diffbio/operators/alignment/profile_hmm.py +350 -0
  45. diffbio/operators/alignment/scoring.py +127 -0
  46. diffbio/operators/alignment/smith_waterman.py +261 -0
  47. diffbio/operators/alignment/soft_msa.py +419 -0
  48. diffbio/operators/assembly/__init__.py +27 -0
  49. diffbio/operators/assembly/gnn_assembly.py +252 -0
  50. diffbio/operators/assembly/metagenomic_binning.py +296 -0
  51. diffbio/operators/crispr/__init__.py +17 -0
  52. diffbio/operators/crispr/guide_scoring.py +269 -0
  53. diffbio/operators/drug_discovery/__init__.py +133 -0
  54. diffbio/operators/drug_discovery/_graph_utils.py +142 -0
  55. diffbio/operators/drug_discovery/admet_predictor.py +285 -0
  56. diffbio/operators/drug_discovery/attentive_fp.py +411 -0
  57. diffbio/operators/drug_discovery/dti.py +261 -0
  58. diffbio/operators/drug_discovery/fingerprint.py +490 -0
  59. diffbio/operators/drug_discovery/maccs_keys.py +267 -0
  60. diffbio/operators/drug_discovery/message_passing.py +200 -0
  61. diffbio/operators/drug_discovery/primitives.py +242 -0
  62. diffbio/operators/drug_discovery/property_predictor.py +163 -0
  63. diffbio/operators/drug_discovery/similarity.py +193 -0
  64. diffbio/operators/epigenomics/__init__.py +35 -0
  65. diffbio/operators/epigenomics/chromatin_state.py +491 -0
  66. diffbio/operators/epigenomics/contextual.py +288 -0
  67. diffbio/operators/epigenomics/fno_peak_calling.py +153 -0
  68. diffbio/operators/epigenomics/peak_calling.py +555 -0
  69. diffbio/operators/foundation_models/__init__.py +119 -0
  70. diffbio/operators/foundation_models/adapters.py +114 -0
  71. diffbio/operators/foundation_models/contracts.py +245 -0
  72. diffbio/operators/foundation_models/embedding_probe.py +83 -0
  73. diffbio/operators/foundation_models/experimental.py +128 -0
  74. diffbio/operators/foundation_models/foundation_model.py +332 -0
  75. diffbio/operators/foundation_models/frozen.py +59 -0
  76. diffbio/operators/foundation_models/precomputed.py +270 -0
  77. diffbio/operators/foundation_models/transformer_encoder.py +564 -0
  78. diffbio/operators/mapping/__init__.py +17 -0
  79. diffbio/operators/mapping/neural_mapper.py +493 -0
  80. diffbio/operators/metabolomics/__init__.py +39 -0
  81. diffbio/operators/metabolomics/spectral_similarity.py +315 -0
  82. diffbio/operators/molecular_dynamics/__init__.py +51 -0
  83. diffbio/operators/molecular_dynamics/force_field.py +265 -0
  84. diffbio/operators/molecular_dynamics/integrator.py +304 -0
  85. diffbio/operators/molecular_dynamics/primitives.py +115 -0
  86. diffbio/operators/multiomics/__init__.py +38 -0
  87. diffbio/operators/multiomics/hic_contact.py +377 -0
  88. diffbio/operators/multiomics/multiomics_vae.py +325 -0
  89. diffbio/operators/multiomics/spatial_deconvolution.py +316 -0
  90. diffbio/operators/multiomics/spatial_gene_detection.py +493 -0
  91. diffbio/operators/normalization/__init__.py +42 -0
  92. diffbio/operators/normalization/embedding.py +222 -0
  93. diffbio/operators/normalization/phate.py +400 -0
  94. diffbio/operators/normalization/umap.py +261 -0
  95. diffbio/operators/normalization/vae_normalizer.py +258 -0
  96. diffbio/operators/population/__init__.py +17 -0
  97. diffbio/operators/population/ancestry_estimation.py +274 -0
  98. diffbio/operators/preprocessing/__init__.py +76 -0
  99. diffbio/operators/preprocessing/adapter_removal.py +311 -0
  100. diffbio/operators/preprocessing/duplicate_filter.py +317 -0
  101. diffbio/operators/preprocessing/error_correction.py +287 -0
  102. diffbio/operators/protein/__init__.py +31 -0
  103. diffbio/operators/protein/secondary_structure.py +509 -0
  104. diffbio/operators/quality_filter.py +128 -0
  105. diffbio/operators/rna_structure/__init__.py +35 -0
  106. diffbio/operators/rna_structure/rna_folding.py +509 -0
  107. diffbio/operators/rnaseq/__init__.py +23 -0
  108. diffbio/operators/rnaseq/motif_discovery.py +251 -0
  109. diffbio/operators/rnaseq/splicing_psi.py +216 -0
  110. diffbio/operators/singlecell/__init__.py +193 -0
  111. diffbio/operators/singlecell/ambient_removal.py +333 -0
  112. diffbio/operators/singlecell/archetypes.py +191 -0
  113. diffbio/operators/singlecell/batch_correction.py +288 -0
  114. diffbio/operators/singlecell/cell_annotation.py +519 -0
  115. diffbio/operators/singlecell/communication.py +704 -0
  116. diffbio/operators/singlecell/differential_distribution.py +243 -0
  117. diffbio/operators/singlecell/doublet_detection.py +657 -0
  118. diffbio/operators/singlecell/downsampling.py +166 -0
  119. diffbio/operators/singlecell/enhanced_batch_correction.py +519 -0
  120. diffbio/operators/singlecell/grn_inference.py +336 -0
  121. diffbio/operators/singlecell/imputation.py +429 -0
  122. diffbio/operators/singlecell/knockdown_filter.py +176 -0
  123. diffbio/operators/singlecell/ot_trajectory.py +277 -0
  124. diffbio/operators/singlecell/simulation.py +444 -0
  125. diffbio/operators/singlecell/sindy_grn.py +247 -0
  126. diffbio/operators/singlecell/soft_clustering.py +211 -0
  127. diffbio/operators/singlecell/spatial_domains.py +677 -0
  128. diffbio/operators/singlecell/switch_de.py +184 -0
  129. diffbio/operators/singlecell/trajectory.py +447 -0
  130. diffbio/operators/singlecell/velocity.py +361 -0
  131. diffbio/operators/statistical/__init__.py +35 -0
  132. diffbio/operators/statistical/em_quantification.py +260 -0
  133. diffbio/operators/statistical/hmm.py +234 -0
  134. diffbio/operators/statistical/nb_glm.py +272 -0
  135. diffbio/operators/variant/__init__.py +64 -0
  136. diffbio/operators/variant/classifier.py +333 -0
  137. diffbio/operators/variant/cnn_classifier.py +255 -0
  138. diffbio/operators/variant/cnv_segmentation.py +678 -0
  139. diffbio/operators/variant/deepvariant_pileup.py +426 -0
  140. diffbio/operators/variant/pileup.py +240 -0
  141. diffbio/operators/variant/quality_recalibration.py +274 -0
  142. diffbio/pipelines/__init__.py +65 -0
  143. diffbio/pipelines/differential_expression.py +279 -0
  144. diffbio/pipelines/enhanced_variant_calling.py +326 -0
  145. diffbio/pipelines/perturbation.py +407 -0
  146. diffbio/pipelines/preprocessing.py +267 -0
  147. diffbio/pipelines/single_cell.py +366 -0
  148. diffbio/pipelines/variant_calling.py +490 -0
  149. diffbio/samplers/__init__.py +9 -0
  150. diffbio/samplers/perturbation_sampler.py +142 -0
  151. diffbio/sequences/__init__.py +34 -0
  152. diffbio/sequences/dna.py +239 -0
  153. diffbio/sources/__init__.py +149 -0
  154. diffbio/sources/_anndata_shared.py +89 -0
  155. diffbio/sources/_batch_iteration.py +37 -0
  156. diffbio/sources/_benchmark_source.py +152 -0
  157. diffbio/sources/_indexed_batch_source.py +38 -0
  158. diffbio/sources/_utils.py +45 -0
  159. diffbio/sources/anndata_interop.py +387 -0
  160. diffbio/sources/anndata_source.py +361 -0
  161. diffbio/sources/archive_ii.py +174 -0
  162. diffbio/sources/balifam.py +207 -0
  163. diffbio/sources/bam.py +265 -0
  164. diffbio/sources/bengrn_ground_truth.py +306 -0
  165. diffbio/sources/contextual_epigenomics.py +242 -0
  166. diffbio/sources/dti.py +359 -0
  167. diffbio/sources/embeddings.py +203 -0
  168. diffbio/sources/encode_peaks.py +223 -0
  169. diffbio/sources/fasta.py +226 -0
  170. diffbio/sources/immune_human.py +172 -0
  171. diffbio/sources/indexed_embeddings.py +128 -0
  172. diffbio/sources/indexed_view.py +191 -0
  173. diffbio/sources/molnet.py +493 -0
  174. diffbio/sources/multiomics.py +279 -0
  175. diffbio/sources/pancreas.py +108 -0
  176. diffbio/sources/perturbation/__init__.py +69 -0
  177. diffbio/sources/perturbation/_types.py +51 -0
  178. diffbio/sources/perturbation/_utils.py +125 -0
  179. diffbio/sources/perturbation/concat_source.py +115 -0
  180. diffbio/sources/perturbation/control_mapping.py +215 -0
  181. diffbio/sources/perturbation/experiment_config.py +261 -0
  182. diffbio/sources/perturbation/h5_metadata_cache.py +218 -0
  183. diffbio/sources/perturbation/output_space.py +52 -0
  184. diffbio/sources/perturbation/perturbation_source.py +513 -0
  185. diffbio/sources/seqfish.py +145 -0
  186. diffbio/sources/sequence_foundation.py +68 -0
  187. diffbio/sources/singlecell_foundation.py +68 -0
  188. diffbio/splitters/__init__.py +63 -0
  189. diffbio/splitters/base.py +251 -0
  190. diffbio/splitters/molecular.py +330 -0
  191. diffbio/splitters/perturbation.py +199 -0
  192. diffbio/splitters/random.py +217 -0
  193. diffbio/splitters/sequence.py +201 -0
  194. diffbio/utils/__init__.py +55 -0
  195. diffbio/utils/dependency_runtime.py +115 -0
  196. diffbio/utils/nn_utils.py +157 -0
  197. diffbio/utils/quality.py +45 -0
  198. diffbio/utils/training.py +585 -0
  199. diffbio-0.1.0.dist-info/METADATA +480 -0
  200. diffbio-0.1.0.dist-info/RECORD +202 -0
  201. diffbio-0.1.0.dist-info/WHEEL +4 -0
  202. 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