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,163 @@
|
|
|
1
|
+
"""Molecular property prediction operator.
|
|
2
|
+
|
|
3
|
+
This module implements a ChemProp-style molecular property predictor
|
|
4
|
+
using message passing neural networks.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import logging
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from artifex.generative_models.core.base import MLP
|
|
12
|
+
from datarax.core.config import OperatorConfig
|
|
13
|
+
from datarax.core.operator import OperatorModule
|
|
14
|
+
from flax import nnx
|
|
15
|
+
|
|
16
|
+
from diffbio.operators.drug_discovery._graph_utils import (
|
|
17
|
+
build_optional_dropout,
|
|
18
|
+
graph_sum_readout,
|
|
19
|
+
initialize_graph_encoder_from_config,
|
|
20
|
+
)
|
|
21
|
+
from diffbio.utils.nn_utils import ARTIFEX_RELU_MLP_KWARGS
|
|
22
|
+
|
|
23
|
+
logger = logging.getLogger(__name__)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass(frozen=True)
|
|
27
|
+
class MolecularPropertyConfig(OperatorConfig):
|
|
28
|
+
"""Configuration for molecular property predictor.
|
|
29
|
+
|
|
30
|
+
Attributes:
|
|
31
|
+
hidden_dim: Hidden dimension for message passing layers.
|
|
32
|
+
num_message_passing_steps: Number of message passing iterations.
|
|
33
|
+
num_output_tasks: Number of prediction tasks (multi-task learning).
|
|
34
|
+
dropout_rate: Dropout rate for regularization.
|
|
35
|
+
in_features: Number of input node features (default: DEFAULT_ATOM_FEATURES=34).
|
|
36
|
+
num_edge_features: Number of edge/bond features.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
hidden_dim: int = 300
|
|
40
|
+
num_message_passing_steps: int = 3
|
|
41
|
+
num_output_tasks: int = 1
|
|
42
|
+
dropout_rate: float = 0.0
|
|
43
|
+
in_features: int = 4 # Default for tests; use DEFAULT_ATOM_FEATURES for real molecules
|
|
44
|
+
num_edge_features: int = 4
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class MolecularPropertyPredictor(OperatorModule):
|
|
48
|
+
"""ChemProp-style molecular property predictor.
|
|
49
|
+
|
|
50
|
+
Implements a directed message passing neural network (D-MPNN) for
|
|
51
|
+
predicting molecular properties from graph representations.
|
|
52
|
+
|
|
53
|
+
The architecture consists of:
|
|
54
|
+
1. Message passing layers to compute atom representations
|
|
55
|
+
2. Graph-level readout via sum pooling
|
|
56
|
+
3. Feed-forward network for property prediction
|
|
57
|
+
|
|
58
|
+
Example:
|
|
59
|
+
```python
|
|
60
|
+
config = MolecularPropertyConfig(hidden_dim=64, num_output_tasks=3)
|
|
61
|
+
predictor = MolecularPropertyPredictor(config, rngs=nnx.Rngs(42))
|
|
62
|
+
data = {
|
|
63
|
+
"node_features": node_features,
|
|
64
|
+
"adjacency": adjacency,
|
|
65
|
+
"node_mask": mask,
|
|
66
|
+
}
|
|
67
|
+
result, state, meta = predictor.apply(data, {}, None)
|
|
68
|
+
predictions = result["predictions"] # shape: (3,)
|
|
69
|
+
```
|
|
70
|
+
"""
|
|
71
|
+
|
|
72
|
+
def __init__(self, config: MolecularPropertyConfig, *, rngs: nnx.Rngs | None = None):
|
|
73
|
+
"""Initialize molecular property predictor.
|
|
74
|
+
|
|
75
|
+
Args:
|
|
76
|
+
config: Predictor configuration.
|
|
77
|
+
rngs: Flax NNX random number generators.
|
|
78
|
+
"""
|
|
79
|
+
super().__init__(config, rngs=rngs)
|
|
80
|
+
|
|
81
|
+
rngs = initialize_graph_encoder_from_config(self, config, rngs=rngs)
|
|
82
|
+
|
|
83
|
+
self.ffn_backbone = MLP(
|
|
84
|
+
hidden_dims=[config.hidden_dim],
|
|
85
|
+
in_features=config.hidden_dim,
|
|
86
|
+
dropout_rate=config.dropout_rate,
|
|
87
|
+
rngs=rngs,
|
|
88
|
+
**ARTIFEX_RELU_MLP_KWARGS,
|
|
89
|
+
)
|
|
90
|
+
self.output_layer = nnx.Linear(config.hidden_dim, config.num_output_tasks, rngs=rngs)
|
|
91
|
+
|
|
92
|
+
self.dropout = build_optional_dropout(config.dropout_rate, rngs=rngs)
|
|
93
|
+
|
|
94
|
+
def apply(
|
|
95
|
+
self,
|
|
96
|
+
data: dict[str, Any],
|
|
97
|
+
state: dict[str, Any],
|
|
98
|
+
metadata: dict[str, Any] | None,
|
|
99
|
+
random_params: Any = None,
|
|
100
|
+
stats: dict[str, Any] | None = None,
|
|
101
|
+
) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any] | None]:
|
|
102
|
+
"""Predict molecular properties from graph representation.
|
|
103
|
+
|
|
104
|
+
Args:
|
|
105
|
+
data: Input data containing:
|
|
106
|
+
- node_features: (num_nodes, num_features) atom features
|
|
107
|
+
- adjacency: (num_nodes, num_nodes) adjacency matrix
|
|
108
|
+
- edge_features: Optional (num_nodes, num_nodes, num_edge_features)
|
|
109
|
+
- node_mask: (num_nodes,) mask for valid nodes
|
|
110
|
+
state: Per-element state (passed through).
|
|
111
|
+
metadata: Optional metadata.
|
|
112
|
+
random_params: Unused random parameters.
|
|
113
|
+
stats: Optional statistics dictionary.
|
|
114
|
+
|
|
115
|
+
Returns:
|
|
116
|
+
Tuple of:
|
|
117
|
+
- data with added "predictions" key
|
|
118
|
+
- unchanged state
|
|
119
|
+
- unchanged metadata
|
|
120
|
+
"""
|
|
121
|
+
graph_repr = graph_sum_readout(data, self.encoder, dropout=self.dropout)
|
|
122
|
+
|
|
123
|
+
# Feed-forward prediction
|
|
124
|
+
ffn_output = self.ffn_backbone(graph_repr)
|
|
125
|
+
if isinstance(ffn_output, tuple):
|
|
126
|
+
raise TypeError("MolecularPropertyPredictor FFN must return a single tensor output.")
|
|
127
|
+
predictions = self.output_layer(ffn_output)
|
|
128
|
+
|
|
129
|
+
result = {
|
|
130
|
+
**data,
|
|
131
|
+
"predictions": predictions,
|
|
132
|
+
"graph_representation": graph_repr,
|
|
133
|
+
}
|
|
134
|
+
|
|
135
|
+
return result, state, metadata
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def create_property_predictor(
|
|
139
|
+
hidden_dim: int = 300,
|
|
140
|
+
num_layers: int = 3,
|
|
141
|
+
num_tasks: int = 1,
|
|
142
|
+
dropout_rate: float = 0.0,
|
|
143
|
+
seed: int = 42,
|
|
144
|
+
) -> MolecularPropertyPredictor:
|
|
145
|
+
"""Create a molecular property predictor.
|
|
146
|
+
|
|
147
|
+
Args:
|
|
148
|
+
hidden_dim: Hidden dimension for message passing.
|
|
149
|
+
num_layers: Number of message passing steps.
|
|
150
|
+
num_tasks: Number of prediction tasks.
|
|
151
|
+
dropout_rate: Dropout rate.
|
|
152
|
+
seed: Random seed.
|
|
153
|
+
|
|
154
|
+
Returns:
|
|
155
|
+
Configured MolecularPropertyPredictor.
|
|
156
|
+
"""
|
|
157
|
+
config = MolecularPropertyConfig(
|
|
158
|
+
hidden_dim=hidden_dim,
|
|
159
|
+
num_message_passing_steps=num_layers,
|
|
160
|
+
num_output_tasks=num_tasks,
|
|
161
|
+
dropout_rate=dropout_rate,
|
|
162
|
+
)
|
|
163
|
+
return MolecularPropertyPredictor(config, rngs=nnx.Rngs(seed))
|
|
@@ -0,0 +1,193 @@
|
|
|
1
|
+
"""Differentiable molecular similarity operator.
|
|
2
|
+
|
|
3
|
+
This module implements differentiable similarity metrics for comparing
|
|
4
|
+
molecular fingerprints, enabling gradient-based optimization of similarity.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import logging
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
import jax.numpy as jnp
|
|
12
|
+
from datarax.core.config import OperatorConfig
|
|
13
|
+
from datarax.core.operator import OperatorModule
|
|
14
|
+
from flax import nnx
|
|
15
|
+
|
|
16
|
+
logger = logging.getLogger(__name__)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class MolecularSimilarityConfig(OperatorConfig):
|
|
21
|
+
"""Configuration for molecular similarity operator.
|
|
22
|
+
|
|
23
|
+
Attributes:
|
|
24
|
+
similarity_type: Type of similarity metric ("tanimoto", "cosine", "dice").
|
|
25
|
+
temperature: Temperature for soft similarity (higher = sharper).
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
similarity_type: str = "tanimoto"
|
|
29
|
+
temperature: float = 1.0
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def tanimoto_similarity(a: jnp.ndarray, b: jnp.ndarray, eps: float = 1e-8) -> jnp.ndarray:
|
|
33
|
+
"""Compute differentiable Tanimoto similarity.
|
|
34
|
+
|
|
35
|
+
For continuous vectors, uses the generalized Tanimoto formula:
|
|
36
|
+
T(a, b) = (a · b) / (|a|² + |b|² - a · b)
|
|
37
|
+
|
|
38
|
+
Args:
|
|
39
|
+
a: First fingerprint vector.
|
|
40
|
+
b: Second fingerprint vector.
|
|
41
|
+
eps: Small constant for numerical stability.
|
|
42
|
+
|
|
43
|
+
Returns:
|
|
44
|
+
Similarity score in [0, 1].
|
|
45
|
+
"""
|
|
46
|
+
dot_product = jnp.sum(a * b)
|
|
47
|
+
norm_a_sq = jnp.sum(a * a)
|
|
48
|
+
norm_b_sq = jnp.sum(b * b)
|
|
49
|
+
|
|
50
|
+
# Generalized Tanimoto for continuous vectors
|
|
51
|
+
similarity = dot_product / (norm_a_sq + norm_b_sq - dot_product + eps)
|
|
52
|
+
|
|
53
|
+
# Clamp to [0, 1] for numerical stability
|
|
54
|
+
return jnp.clip(similarity, 0.0, 1.0)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def cosine_similarity(a: jnp.ndarray, b: jnp.ndarray, eps: float = 1e-8) -> jnp.ndarray:
|
|
58
|
+
"""Compute cosine similarity.
|
|
59
|
+
|
|
60
|
+
Args:
|
|
61
|
+
a: First vector.
|
|
62
|
+
b: Second vector.
|
|
63
|
+
eps: Small constant for numerical stability.
|
|
64
|
+
|
|
65
|
+
Returns:
|
|
66
|
+
Similarity score in [-1, 1].
|
|
67
|
+
"""
|
|
68
|
+
dot_product = jnp.sum(a * b)
|
|
69
|
+
norm_a = jnp.linalg.norm(a)
|
|
70
|
+
norm_b = jnp.linalg.norm(b)
|
|
71
|
+
|
|
72
|
+
return dot_product / (norm_a * norm_b + eps)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def dice_similarity(a: jnp.ndarray, b: jnp.ndarray, eps: float = 1e-8) -> jnp.ndarray:
|
|
76
|
+
"""Compute Dice similarity coefficient.
|
|
77
|
+
|
|
78
|
+
For continuous vectors:
|
|
79
|
+
Dice(a, b) = 2 * (a · b) / (|a|² + |b|²)
|
|
80
|
+
|
|
81
|
+
Args:
|
|
82
|
+
a: First vector.
|
|
83
|
+
b: Second vector.
|
|
84
|
+
eps: Small constant for numerical stability.
|
|
85
|
+
|
|
86
|
+
Returns:
|
|
87
|
+
Similarity score in [0, 1].
|
|
88
|
+
"""
|
|
89
|
+
dot_product = jnp.sum(a * b)
|
|
90
|
+
norm_a_sq = jnp.sum(a * a)
|
|
91
|
+
norm_b_sq = jnp.sum(b * b)
|
|
92
|
+
|
|
93
|
+
similarity = 2 * dot_product / (norm_a_sq + norm_b_sq + eps)
|
|
94
|
+
|
|
95
|
+
return jnp.clip(similarity, 0.0, 1.0)
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
class MolecularSimilarityOperator(OperatorModule):
|
|
99
|
+
"""Differentiable molecular similarity operator.
|
|
100
|
+
|
|
101
|
+
Computes similarity between molecular fingerprints using various
|
|
102
|
+
differentiable metrics. Supports Tanimoto, cosine, and Dice similarity.
|
|
103
|
+
|
|
104
|
+
Example:
|
|
105
|
+
```python
|
|
106
|
+
config = MolecularSimilarityConfig(similarity_type="tanimoto")
|
|
107
|
+
sim_op = MolecularSimilarityOperator(config, rngs=nnx.Rngs(42))
|
|
108
|
+
data = {"fingerprint_a": fp1, "fingerprint_b": fp2}
|
|
109
|
+
result, _, _ = sim_op.apply(data, {}, None)
|
|
110
|
+
similarity = result["similarity"] # scalar in [0, 1]
|
|
111
|
+
```
|
|
112
|
+
"""
|
|
113
|
+
|
|
114
|
+
def __init__(self, config: MolecularSimilarityConfig, *, rngs: nnx.Rngs | None = None):
|
|
115
|
+
"""Initialize similarity operator.
|
|
116
|
+
|
|
117
|
+
Args:
|
|
118
|
+
config: Similarity configuration.
|
|
119
|
+
rngs: Flax NNX random number generators.
|
|
120
|
+
"""
|
|
121
|
+
super().__init__(config, rngs=rngs)
|
|
122
|
+
self.config: MolecularSimilarityConfig = config
|
|
123
|
+
|
|
124
|
+
# Fix: wrap _unique_id as static for jax.grad compatibility
|
|
125
|
+
# (datarax stores it as plain int which causes gradient errors)
|
|
126
|
+
self._unique_id = nnx.static(self._unique_id)
|
|
127
|
+
|
|
128
|
+
# Select similarity function
|
|
129
|
+
if config.similarity_type == "tanimoto":
|
|
130
|
+
self._similarity_fn = tanimoto_similarity
|
|
131
|
+
elif config.similarity_type == "cosine":
|
|
132
|
+
self._similarity_fn = cosine_similarity
|
|
133
|
+
elif config.similarity_type == "dice":
|
|
134
|
+
self._similarity_fn = dice_similarity
|
|
135
|
+
else:
|
|
136
|
+
raise ValueError(f"Unknown similarity type: {config.similarity_type}")
|
|
137
|
+
|
|
138
|
+
def apply(
|
|
139
|
+
self,
|
|
140
|
+
data: dict[str, Any],
|
|
141
|
+
state: dict[str, Any],
|
|
142
|
+
metadata: dict[str, Any] | None,
|
|
143
|
+
random_params: Any = None,
|
|
144
|
+
stats: dict[str, Any] | None = None,
|
|
145
|
+
) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any] | None]:
|
|
146
|
+
"""Compute similarity between two fingerprints.
|
|
147
|
+
|
|
148
|
+
Args:
|
|
149
|
+
data: Input data containing:
|
|
150
|
+
- fingerprint_a: First fingerprint vector
|
|
151
|
+
- fingerprint_b: Second fingerprint vector
|
|
152
|
+
state: Per-element state (passed through).
|
|
153
|
+
metadata: Optional metadata.
|
|
154
|
+
random_params: Unused random parameters.
|
|
155
|
+
stats: Optional statistics dictionary.
|
|
156
|
+
|
|
157
|
+
Returns:
|
|
158
|
+
Tuple of:
|
|
159
|
+
- data with added "similarity" key
|
|
160
|
+
- unchanged state
|
|
161
|
+
- unchanged metadata
|
|
162
|
+
"""
|
|
163
|
+
fp_a = data["fingerprint_a"]
|
|
164
|
+
fp_b = data["fingerprint_b"]
|
|
165
|
+
|
|
166
|
+
similarity = self._similarity_fn(fp_a, fp_b)
|
|
167
|
+
|
|
168
|
+
result = {
|
|
169
|
+
**data,
|
|
170
|
+
"similarity": similarity,
|
|
171
|
+
}
|
|
172
|
+
|
|
173
|
+
return result, state, metadata
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def create_similarity_operator(
|
|
177
|
+
similarity_type: str = "tanimoto",
|
|
178
|
+
temperature: float = 1.0,
|
|
179
|
+
) -> MolecularSimilarityOperator:
|
|
180
|
+
"""Create a molecular similarity operator.
|
|
181
|
+
|
|
182
|
+
Args:
|
|
183
|
+
similarity_type: Type of similarity ("tanimoto", "cosine", "dice").
|
|
184
|
+
temperature: Temperature parameter.
|
|
185
|
+
|
|
186
|
+
Returns:
|
|
187
|
+
Configured MolecularSimilarityOperator.
|
|
188
|
+
"""
|
|
189
|
+
config = MolecularSimilarityConfig(
|
|
190
|
+
similarity_type=similarity_type,
|
|
191
|
+
temperature=temperature,
|
|
192
|
+
)
|
|
193
|
+
return MolecularSimilarityOperator(config, rngs=nnx.Rngs(42))
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""Epigenomics operators for differentiable ChIP-seq and ATAC-seq analysis.
|
|
2
|
+
|
|
3
|
+
This module provides differentiable operators for epigenomic data analysis,
|
|
4
|
+
including peak calling and chromatin state annotation.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from diffbio.operators.epigenomics.chromatin_state import (
|
|
8
|
+
ChromatinStateAnnotator,
|
|
9
|
+
ChromatinStateConfig,
|
|
10
|
+
)
|
|
11
|
+
from diffbio.operators.epigenomics.contextual import (
|
|
12
|
+
ContextualEpigenomicsConfig,
|
|
13
|
+
ContextualEpigenomicsOperator,
|
|
14
|
+
compute_contextual_epigenomics_loss,
|
|
15
|
+
)
|
|
16
|
+
from diffbio.operators.epigenomics.fno_peak_calling import (
|
|
17
|
+
FNOPeakCaller,
|
|
18
|
+
FNOPeakCallerConfig,
|
|
19
|
+
)
|
|
20
|
+
from diffbio.operators.epigenomics.peak_calling import (
|
|
21
|
+
DifferentiablePeakCaller,
|
|
22
|
+
PeakCallerConfig,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
__all__ = [
|
|
26
|
+
"ChromatinStateAnnotator",
|
|
27
|
+
"ChromatinStateConfig",
|
|
28
|
+
"ContextualEpigenomicsConfig",
|
|
29
|
+
"ContextualEpigenomicsOperator",
|
|
30
|
+
"DifferentiablePeakCaller",
|
|
31
|
+
"FNOPeakCaller",
|
|
32
|
+
"FNOPeakCallerConfig",
|
|
33
|
+
"PeakCallerConfig",
|
|
34
|
+
"compute_contextual_epigenomics_loss",
|
|
35
|
+
]
|