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