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,564 @@
1
+ """Transformer-based sequence encoder for DNA/RNA foundation models.
2
+
3
+ This module provides a differentiable transformer encoder following
4
+ DNABERT/RNA-FM architecture patterns. The encoder converts one-hot
5
+ encoded nucleotide sequences into dense embeddings suitable for
6
+ downstream bioinformatics tasks.
7
+
8
+ Key features:
9
+
10
+ - Multi-head self-attention for capturing sequence dependencies
11
+ - Sinusoidal positional encoding for position awareness
12
+ - Configurable architecture (layers, heads, dimensions)
13
+ - Multiple pooling strategies (mean, CLS token)
14
+ - Fully differentiable for gradient-based optimization
15
+
16
+ References:
17
+ - DNABERT: Ji et al. (2021) Bioinformatics
18
+ - RNA-FM: Chen et al. (2022) Nature Methods
19
+ """
20
+
21
+ import logging
22
+ from dataclasses import dataclass
23
+ from typing import Any, Literal
24
+
25
+ import jax
26
+ import jax.numpy as jnp
27
+ from artifex.generative_models.core.layers import TransformerEncoder
28
+ from flax import nnx
29
+ from jaxtyping import Array, Float, PyTree
30
+
31
+ from diffbio.core.base_operators import SequenceOperator
32
+ from diffbio.operators._transformer_validation import TransformerEncoderShapeValidationMixin
33
+ from diffbio.operators.foundation_models.contracts import (
34
+ FoundationEmbeddingMixin,
35
+ FoundationEmbeddingOperatorConfig,
36
+ FoundationModelKind,
37
+ PoolingStrategy,
38
+ register_foundation_model,
39
+ )
40
+
41
+ logger = logging.getLogger(__name__)
42
+
43
+
44
+ @dataclass(frozen=True)
45
+ class _TransformerArchitectureConfig:
46
+ """Transformer depth and width configuration."""
47
+
48
+ hidden_dim: int = 256
49
+ num_layers: int = 4
50
+ num_heads: int = 4
51
+ intermediate_dim: int = 1024
52
+ max_length: int = 512
53
+
54
+
55
+ @dataclass(frozen=True)
56
+ class _TransformerInputConfig:
57
+ """Sequence input encoding configuration."""
58
+
59
+ alphabet_size: int = 4
60
+ input_embedding_type: Literal["linear", "token_embedding"] = "linear"
61
+ vocab_size: int | None = None
62
+
63
+
64
+ @dataclass(frozen=True)
65
+ class _TransformerOutputConfig:
66
+ """Transformer output and artifact configuration."""
67
+
68
+ dropout_rate: float = 0.1
69
+ pooling: Literal["mean", "cls"] = "mean"
70
+ artifact_id: str = "diffbio.transformer_sequence_encoder"
71
+ preprocessing_version: str = "one_hot_v1"
72
+
73
+
74
+ @dataclass(frozen=True)
75
+ class TransformerSequenceEncoderConfig(
76
+ _TransformerArchitectureConfig,
77
+ _TransformerInputConfig,
78
+ _TransformerOutputConfig,
79
+ TransformerEncoderShapeValidationMixin,
80
+ FoundationEmbeddingOperatorConfig,
81
+ ):
82
+ """Configuration for TransformerSequenceEncoder."""
83
+
84
+ def __post_init__(self) -> None:
85
+ """Validate the transformer encoder configuration."""
86
+ super().__post_init__()
87
+ if self.alphabet_size <= 0:
88
+ raise ValueError("alphabet_size must be positive.")
89
+
90
+ try:
91
+ PoolingStrategy(self.pooling)
92
+ except ValueError as exc:
93
+ raise ValueError("pooling must be 'mean' or 'cls'.") from exc
94
+
95
+ if self.input_embedding_type not in ("linear", "token_embedding"):
96
+ raise ValueError("input_embedding_type must be 'linear' or 'token_embedding'.")
97
+ if self.input_embedding_type == "token_embedding":
98
+ if self.vocab_size is None:
99
+ raise ValueError(
100
+ "vocab_size must be specified when input_embedding_type is 'token_embedding'"
101
+ )
102
+ if self.vocab_size <= 0:
103
+ raise ValueError("vocab_size must be positive.")
104
+ elif self.vocab_size is not None and self.vocab_size <= 0:
105
+ raise ValueError("vocab_size must be positive when provided.")
106
+
107
+
108
+ class TransformerSequenceEncoder(FoundationEmbeddingMixin, SequenceOperator):
109
+ """Transformer-based encoder for DNA/RNA sequences.
110
+
111
+ This operator implements a BERT-style transformer encoder that
112
+ converts nucleotide sequences into dense embeddings. The architecture
113
+ follows DNABERT and RNA-FM patterns.
114
+
115
+ Uses artifex's TransformerEncoder for the core transformer layers,
116
+ following the DRY principle.
117
+
118
+ Supports two input embedding modes:
119
+
120
+ - "linear" (default): Projects one-hot encoded input (seq_len, alphabet_size)
121
+ via nnx.Linear. This is the standard mode for continuous one-hot input.
122
+ - "token_embedding": Embeds integer token IDs (seq_len,) via nnx.Embed.
123
+ Useful for gene-token foundation models and tokenized input.
124
+
125
+ The encoder produces:
126
+
127
+ - Global sequence embedding via mean pooling or CLS token
128
+ - Per-position embeddings for fine-grained analysis
129
+
130
+ Args:
131
+ config: TransformerSequenceEncoderConfig with model parameters.
132
+ rngs: Flax NNX random number generators.
133
+ name: Optional operator name.
134
+
135
+ Example:
136
+ ```python
137
+ config = TransformerSequenceEncoderConfig(hidden_dim=256)
138
+ encoder = TransformerSequenceEncoder(config, rngs=nnx.Rngs(42))
139
+ data = {"sequence": one_hot_sequence}
140
+ result, state, meta = encoder.apply(data, {}, None)
141
+ embeddings = result["embeddings"]
142
+ ```
143
+ """
144
+
145
+ foundation_model_kind = FoundationModelKind.SEQUENCE_TRANSFORMER
146
+
147
+ def __init__(
148
+ self,
149
+ config: TransformerSequenceEncoderConfig,
150
+ *,
151
+ rngs: nnx.Rngs | None = None,
152
+ name: str | None = None,
153
+ ):
154
+ """Initialize the transformer encoder.
155
+
156
+ Args:
157
+ config: Encoder configuration.
158
+ rngs: Random number generators for initialization.
159
+ name: Optional operator name.
160
+ """
161
+ super().__init__(config, rngs=rngs, name=name)
162
+
163
+ if rngs is None:
164
+ rngs = nnx.Rngs(0)
165
+
166
+ # Ensure dropout stream exists for artifex transformer
167
+ if config.dropout_rate > 0 and "dropout" not in rngs:
168
+ rngs = nnx.Rngs(params=rngs.params(), dropout=jax.random.key(1))
169
+
170
+ # Input projection: alphabet_size -> hidden_dim (or token embedding)
171
+ if config.input_embedding_type == "token_embedding":
172
+ assert config.vocab_size is not None
173
+ self.input_projection = nnx.Embed(
174
+ num_embeddings=config.vocab_size,
175
+ features=config.hidden_dim,
176
+ rngs=rngs,
177
+ )
178
+ else:
179
+ self.input_projection = nnx.Linear(
180
+ config.alphabet_size,
181
+ config.hidden_dim,
182
+ rngs=rngs,
183
+ )
184
+
185
+ # CLS token embedding (learnable)
186
+ self.cls_token = nnx.Param(jax.random.normal(rngs.params(), (config.hidden_dim,)) * 0.02)
187
+
188
+ # Compute MLP ratio from intermediate_dim
189
+ mlp_ratio = config.intermediate_dim / config.hidden_dim
190
+
191
+ # Use artifex's TransformerEncoder (DRY principle)
192
+ self.transformer = TransformerEncoder(
193
+ num_layers=config.num_layers,
194
+ hidden_dim=config.hidden_dim,
195
+ num_heads=config.num_heads,
196
+ mlp_ratio=mlp_ratio,
197
+ dropout_rate=config.dropout_rate,
198
+ attention_dropout_rate=0.0,
199
+ max_len=config.max_length + 1, # +1 for CLS token
200
+ pos_encoding_type="sinusoidal",
201
+ rngs=rngs,
202
+ )
203
+
204
+ def foundation_pooling_strategy(self) -> PoolingStrategy:
205
+ """Return the pooling strategy for the global sequence embedding."""
206
+ return PoolingStrategy(self.config.pooling)
207
+
208
+ def get_positional_encoding(
209
+ self,
210
+ seq_len: int,
211
+ ) -> Float[Array, "seq_len hidden_dim"]:
212
+ """Generate sinusoidal positional encoding.
213
+
214
+ This is provided for compatibility but the transformer uses
215
+ internal positional encoding.
216
+
217
+ Args:
218
+ seq_len: Sequence length.
219
+
220
+ Returns:
221
+ Positional encoding matrix.
222
+ """
223
+ hidden_dim = self.config.hidden_dim
224
+ position = jnp.arange(seq_len)[:, None]
225
+ div_term = jnp.exp(jnp.arange(0, hidden_dim, 2) * -(jnp.log(10000.0) / hidden_dim))
226
+
227
+ pe = jnp.zeros((seq_len, hidden_dim))
228
+ pe = pe.at[:, 0::2].set(jnp.sin(position * div_term))
229
+ pe = pe.at[:, 1::2].set(jnp.cos(position * div_term))
230
+
231
+ return pe
232
+
233
+ def _encode_single(
234
+ self,
235
+ sequence: Array,
236
+ mask: Float[Array, "seq_len"] | None = None,
237
+ ) -> tuple[Float[Array, "hidden_dim"], Float[Array, "seq_len hidden_dim"]]:
238
+ """Encode a single sequence.
239
+
240
+ Args:
241
+ sequence: Input sequence. One-hot encoded (seq_len, alphabet_size)
242
+ for linear mode, or integer token IDs (seq_len,) for token
243
+ embedding mode.
244
+ mask: Optional attention mask.
245
+
246
+ Returns:
247
+ Tuple of (global_embedding, token_embeddings).
248
+ """
249
+ # Project input to hidden dimension
250
+ hidden = self.input_projection(sequence)
251
+
252
+ # Add batch dimension for transformer (expects [batch, seq, hidden])
253
+ hidden = hidden[None, :, :] # (1, seq_len, hidden_dim)
254
+
255
+ # Prepend CLS token for CLS pooling
256
+ if self.config.pooling == "cls":
257
+ cls_token = self.cls_token[...][None, None, :] # (1, 1, hidden_dim)
258
+ hidden = jnp.concatenate([cls_token, hidden], axis=1)
259
+
260
+ # Extend mask if provided
261
+ if mask is not None:
262
+ mask = jnp.concatenate([jnp.ones(1), mask], axis=0)
263
+
264
+ # Ensure mask has batch dimension for artifex transformer
265
+ if mask is not None:
266
+ mask = mask[None, :] # Add batch dim: (1, seq_len)
267
+
268
+ # Apply transformer (deterministic=True for no dropout)
269
+ hidden = self.transformer(hidden, mask=mask, deterministic=True)
270
+
271
+ # Remove batch dimension
272
+ hidden = hidden[0] # (seq_len, hidden_dim)
273
+
274
+ # Extract embeddings based on pooling strategy
275
+ if self.config.pooling == "cls":
276
+ # Use CLS token (first position)
277
+ global_embedding = hidden[0]
278
+ position_embeddings = hidden[1:] # Remove CLS token
279
+ else:
280
+ # Mean pooling
281
+ if mask is not None:
282
+ # Mask is (1, seq_len), get the 1D version
283
+ mask_1d = mask[0]
284
+ mask_expanded = mask_1d[:, None]
285
+ masked_hidden = hidden * mask_expanded
286
+ global_embedding = jnp.sum(masked_hidden, axis=0) / (jnp.sum(mask_1d) + 1e-9)
287
+ else:
288
+ global_embedding = jnp.mean(hidden, axis=0)
289
+ position_embeddings = hidden
290
+
291
+ return global_embedding, position_embeddings
292
+
293
+ def _encode_batch(
294
+ self,
295
+ sequences: Array,
296
+ masks: Float[Array, "batch seq_len"] | None = None,
297
+ ) -> tuple[
298
+ Float[Array, "batch hidden_dim"],
299
+ Float[Array, "batch seq_len hidden_dim"],
300
+ ]:
301
+ """Encode a batch of sequences.
302
+
303
+ Args:
304
+ sequences: Batch of input sequences. One-hot encoded
305
+ (batch, seq_len, alphabet_size) for linear mode, or integer
306
+ token IDs (batch, seq_len) for token embedding mode.
307
+ masks: Optional attention masks.
308
+
309
+ Returns:
310
+ Tuple of (global_embeddings, token_embeddings).
311
+ """
312
+ batch_size = sequences.shape[0]
313
+
314
+ # Project input to hidden dimension
315
+ hidden = jax.vmap(self.input_projection)(sequences)
316
+
317
+ # Prepend CLS token for CLS pooling
318
+ if self.config.pooling == "cls":
319
+ cls_token = self.cls_token[...][None, None, :] # (1, 1, hidden_dim)
320
+ cls_tokens = jnp.broadcast_to(cls_token, (batch_size, 1, self.config.hidden_dim))
321
+ hidden = jnp.concatenate([cls_tokens, hidden], axis=1)
322
+
323
+ # Extend masks if provided
324
+ if masks is not None:
325
+ mask_prefix = jnp.ones((batch_size, 1))
326
+ masks = jnp.concatenate([mask_prefix, masks], axis=1)
327
+
328
+ # Apply transformer
329
+ hidden = self.transformer(hidden, mask=masks, deterministic=True)
330
+
331
+ # Extract embeddings based on pooling strategy
332
+ if self.config.pooling == "cls":
333
+ global_embeddings = hidden[:, 0]
334
+ position_embeddings = hidden[:, 1:]
335
+ else:
336
+ if masks is not None:
337
+ mask_expanded = masks[:, :, None]
338
+ masked_hidden = hidden * mask_expanded
339
+ global_embeddings = jnp.sum(masked_hidden, axis=1) / (
340
+ jnp.sum(masks, axis=1, keepdims=True) + 1e-9
341
+ )
342
+ else:
343
+ global_embeddings = jnp.mean(hidden, axis=1)
344
+ position_embeddings = hidden
345
+
346
+ return global_embeddings, position_embeddings
347
+
348
+ def apply(
349
+ self,
350
+ data: PyTree,
351
+ state: PyTree,
352
+ metadata: dict[str, Any] | None,
353
+ random_params: Any = None,
354
+ stats: dict[str, Any] | None = None,
355
+ ) -> tuple[PyTree, PyTree, dict[str, Any] | None]:
356
+ """Apply transformer encoding to sequence data.
357
+
358
+ This method encodes DNA/RNA sequences into dense embeddings using
359
+ a transformer encoder architecture.
360
+
361
+ Input shape depends on ``input_embedding_type``:
362
+
363
+ - "linear": one-hot ``(seq_len, alphabet_size)`` or
364
+ ``(batch, seq_len, alphabet_size)``
365
+ - "token_embedding": integer token IDs ``(seq_len,)`` or
366
+ ``(batch, seq_len)``
367
+
368
+ Args:
369
+ data: Dictionary containing:
370
+ - "sequence": Encoded sequence(s) (see above for shapes)
371
+ - "attention_mask": Optional mask (seq_len,) or (batch, seq_len)
372
+ state: Element state (passed through unchanged)
373
+ metadata: Element metadata (passed through unchanged)
374
+ random_params: Not used
375
+ stats: Not used
376
+
377
+ Returns:
378
+ Tuple of (transformed_data, state, metadata):
379
+ - transformed_data contains:
380
+
381
+ - All original keys from data
382
+ - "embeddings": Global sequence embedding
383
+ - "token_embeddings": Per-position hidden states
384
+ - "foundation_model": Canonical artifact metadata
385
+ - state is passed through unchanged
386
+ - metadata is passed through unchanged
387
+ """
388
+ del random_params, stats # Unused
389
+
390
+ sequence = data["sequence"]
391
+ mask = data.get("attention_mask", None)
392
+
393
+ is_token_mode = self.config.input_embedding_type == "token_embedding"
394
+
395
+ # Determine single vs batch based on input dimensionality:
396
+ # - token mode: single=(seq_len,) ndim=1, batch=(batch, seq_len) ndim=2
397
+ # - linear mode: single=(seq_len, alphabet) ndim=2, batch=(batch, seq_len, alphabet) ndim=3
398
+ single_ndim = 1 if is_token_mode else 2
399
+
400
+ if sequence.ndim == single_ndim:
401
+ embeddings, token_embeddings = self._encode_single(sequence, mask)
402
+ else:
403
+ embeddings, token_embeddings = self._encode_batch(sequence, mask)
404
+
405
+ transformed_data = self.foundation_result(
406
+ data,
407
+ embeddings,
408
+ token_embeddings=token_embeddings,
409
+ )
410
+
411
+ return transformed_data, state, metadata
412
+
413
+
414
+ def _create_sequence_encoder(
415
+ alphabet_size: int,
416
+ hidden_dim: int = 256,
417
+ num_layers: int = 4,
418
+ num_heads: int = 4,
419
+ intermediate_dim: int | None = None,
420
+ max_length: int = 512,
421
+ dropout_rate: float = 0.1,
422
+ pooling: Literal["mean", "cls"] = "mean",
423
+ *,
424
+ rngs: nnx.Rngs | None = None,
425
+ ) -> TransformerSequenceEncoder:
426
+ """Create a transformer sequence encoder with given alphabet size.
427
+
428
+ Args:
429
+ alphabet_size: Size of input alphabet (e.g., 4 for DNA/RNA).
430
+ hidden_dim: Dimension of hidden states.
431
+ num_layers: Number of transformer layers.
432
+ num_heads: Number of attention heads.
433
+ intermediate_dim: FFN intermediate dimension (default: 4 * hidden_dim).
434
+ max_length: Maximum sequence length.
435
+ dropout_rate: Dropout rate.
436
+ pooling: Pooling strategy.
437
+ rngs: Random number generators.
438
+
439
+ Returns:
440
+ Configured TransformerSequenceEncoder.
441
+ """
442
+ if intermediate_dim is None:
443
+ intermediate_dim = 4 * hidden_dim
444
+
445
+ if rngs is None:
446
+ rngs = nnx.Rngs(0)
447
+
448
+ config = TransformerSequenceEncoderConfig(
449
+ hidden_dim=hidden_dim,
450
+ num_layers=num_layers,
451
+ num_heads=num_heads,
452
+ intermediate_dim=intermediate_dim,
453
+ max_length=max_length,
454
+ alphabet_size=alphabet_size,
455
+ dropout_rate=dropout_rate,
456
+ pooling=pooling,
457
+ )
458
+
459
+ return TransformerSequenceEncoder(config, rngs=rngs)
460
+
461
+
462
+ def create_dna_encoder(
463
+ hidden_dim: int = 256,
464
+ num_layers: int = 4,
465
+ num_heads: int = 4,
466
+ intermediate_dim: int | None = None,
467
+ max_length: int = 512,
468
+ dropout_rate: float = 0.1,
469
+ pooling: Literal["mean", "cls"] = "mean",
470
+ *,
471
+ rngs: nnx.Rngs | None = None,
472
+ ) -> TransformerSequenceEncoder:
473
+ """Create a transformer encoder for DNA sequences.
474
+
475
+ Factory function for creating a DNA sequence encoder with
476
+ sensible defaults for DNA processing.
477
+
478
+ Args:
479
+ hidden_dim: Dimension of hidden states.
480
+ num_layers: Number of transformer layers.
481
+ num_heads: Number of attention heads.
482
+ intermediate_dim: FFN intermediate dimension (default: 4 * hidden_dim).
483
+ max_length: Maximum sequence length.
484
+ dropout_rate: Dropout rate.
485
+ pooling: Pooling strategy.
486
+ rngs: Random number generators.
487
+
488
+ Returns:
489
+ Configured TransformerSequenceEncoder for DNA.
490
+
491
+ Example:
492
+ ```python
493
+ encoder = create_dna_encoder(hidden_dim=256, num_layers=6)
494
+ data = {"sequence": dna_one_hot}
495
+ result, _, _ = encoder.apply(data, {}, None)
496
+ embeddings = result["embeddings"]
497
+ ```
498
+ """
499
+ return _create_sequence_encoder(
500
+ alphabet_size=4, # A, C, G, T
501
+ hidden_dim=hidden_dim,
502
+ num_layers=num_layers,
503
+ num_heads=num_heads,
504
+ intermediate_dim=intermediate_dim,
505
+ max_length=max_length,
506
+ dropout_rate=dropout_rate,
507
+ pooling=pooling,
508
+ rngs=rngs,
509
+ )
510
+
511
+
512
+ register_foundation_model(
513
+ FoundationModelKind.SEQUENCE_TRANSFORMER,
514
+ TransformerSequenceEncoder,
515
+ )
516
+
517
+
518
+ def create_rna_encoder(
519
+ hidden_dim: int = 256,
520
+ num_layers: int = 4,
521
+ num_heads: int = 4,
522
+ intermediate_dim: int | None = None,
523
+ max_length: int = 512,
524
+ dropout_rate: float = 0.1,
525
+ pooling: Literal["mean", "cls"] = "mean",
526
+ *,
527
+ rngs: nnx.Rngs | None = None,
528
+ ) -> TransformerSequenceEncoder:
529
+ """Create a transformer encoder for RNA sequences.
530
+
531
+ Factory function for creating an RNA sequence encoder with
532
+ sensible defaults for RNA processing.
533
+
534
+ Args:
535
+ hidden_dim: Dimension of hidden states.
536
+ num_layers: Number of transformer layers.
537
+ num_heads: Number of attention heads.
538
+ intermediate_dim: FFN intermediate dimension (default: 4 * hidden_dim).
539
+ max_length: Maximum sequence length.
540
+ dropout_rate: Dropout rate.
541
+ pooling: Pooling strategy.
542
+ rngs: Random number generators.
543
+
544
+ Returns:
545
+ Configured TransformerSequenceEncoder for RNA.
546
+
547
+ Example:
548
+ ```python
549
+ encoder = create_rna_encoder(hidden_dim=640, num_layers=12)
550
+ data = {"sequence": rna_one_hot}
551
+ result, _, _ = encoder.apply(data, {}, None)
552
+ ```
553
+ """
554
+ return _create_sequence_encoder(
555
+ alphabet_size=4, # A, C, G, U
556
+ hidden_dim=hidden_dim,
557
+ num_layers=num_layers,
558
+ num_heads=num_heads,
559
+ intermediate_dim=intermediate_dim,
560
+ max_length=max_length,
561
+ dropout_rate=dropout_rate,
562
+ pooling=pooling,
563
+ rngs=rngs,
564
+ )
@@ -0,0 +1,17 @@
1
+ """Mapping operators for differentiable read alignment.
2
+
3
+ This module provides neural network-based approaches to read mapping
4
+ that enable gradient flow through the mapping process.
5
+
6
+ - NeuralReadMapper: Cross-attention based soft read mapping
7
+ """
8
+
9
+ from diffbio.operators.mapping.neural_mapper import (
10
+ NeuralReadMapper,
11
+ NeuralReadMapperConfig,
12
+ )
13
+
14
+ __all__ = [
15
+ "NeuralReadMapper",
16
+ "NeuralReadMapperConfig",
17
+ ]