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,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
|
+
]
|