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,260 @@
1
+ """Type definitions and protocols for DiffBio.
2
+
3
+ This module provides type aliases, TypedDicts, and protocols that define
4
+ the expected interfaces and data structures across the DiffBio codebase.
5
+ """
6
+
7
+ from typing import Any, Protocol, TypedDict, runtime_checkable
8
+
9
+ from jaxtyping import Array, Float
10
+
11
+ # =============================================================================
12
+ # Type Aliases for Scalar Values
13
+ # =============================================================================
14
+
15
+ Temperature = float
16
+ """Temperature parameter for soft operations. Must be > 0."""
17
+
18
+ Probability = float
19
+ """Probability value in range [0, 1]."""
20
+
21
+ LogProbability = float
22
+ """Log probability value in range (-inf, 0]."""
23
+
24
+ # =============================================================================
25
+ # Type Aliases for Arrays
26
+ # =============================================================================
27
+
28
+ SequenceArray = Float[Array, "length alphabet"]
29
+ """One-hot encoded sequence of shape (length, alphabet_size)."""
30
+
31
+ BatchArray = Float[Array, "batch ..."]
32
+ """Batched array with batch dimension first."""
33
+
34
+ ProbabilityArray = Float[Array, "..."]
35
+ """Array of probability values, each in [0, 1]."""
36
+
37
+ ScoreMatrix = Float[Array, "alphabet alphabet"]
38
+ """Scoring matrix for sequence alignment."""
39
+
40
+ AlignmentMatrix = Float[Array, "len1_plus1 len2_plus1"]
41
+ """Dynamic programming matrix for alignment."""
42
+
43
+ PositionWeightMatrix = Float[Array, "length alphabet"]
44
+ """Position weight matrix for motif representation."""
45
+
46
+
47
+ # =============================================================================
48
+ # TypedDicts for Data Structures
49
+ # =============================================================================
50
+
51
+
52
+ class SequenceData(TypedDict, total=False):
53
+ """Data dictionary for sequence data.
54
+
55
+ Required:
56
+ sequence: One-hot encoded sequence.
57
+
58
+ Optional:
59
+ quality_scores: Phred quality scores.
60
+ mask: Boolean mask for valid positions.
61
+ """
62
+
63
+ sequence: Array
64
+ quality_scores: Array
65
+ mask: Array
66
+
67
+
68
+ class AlignmentResultData(TypedDict, total=False):
69
+ """Data dictionary for alignment results.
70
+
71
+ Required:
72
+ score: Alignment score.
73
+ alignment_matrix: DP matrix.
74
+
75
+ Optional:
76
+ soft_alignment: Soft position correspondences.
77
+ traceback: Hard alignment path.
78
+ """
79
+
80
+ score: Array
81
+ alignment_matrix: Array
82
+ soft_alignment: Array
83
+ traceback: Array
84
+
85
+
86
+ class VariantData(TypedDict, total=False):
87
+ """Data dictionary for variant calling results.
88
+
89
+ Required:
90
+ logits: Classification logits.
91
+
92
+ Optional:
93
+ probabilities: Softmax probabilities.
94
+ pileup: Pileup representation.
95
+ coverage: Coverage at each position.
96
+ """
97
+
98
+ logits: Array
99
+ probabilities: Array
100
+ pileup: Array
101
+ coverage: Array
102
+
103
+
104
+ class LatentData(TypedDict, total=False):
105
+ """Data dictionary for VAE latent representations.
106
+
107
+ Required:
108
+ z: Sampled latent representation.
109
+
110
+ Optional:
111
+ mean: Mean of latent distribution.
112
+ log_var: Log variance of latent distribution.
113
+ """
114
+
115
+ z: Array
116
+ mean: Array
117
+ log_var: Array
118
+
119
+
120
+ class GraphData(TypedDict, total=False):
121
+ """Data dictionary for graph-structured data.
122
+
123
+ Required:
124
+ node_features: Node feature matrix.
125
+ edge_index: Edge indices (2, num_edges).
126
+
127
+ Optional:
128
+ edge_features: Edge feature matrix.
129
+ batch: Batch assignment for nodes.
130
+ """
131
+
132
+ node_features: Array
133
+ edge_index: Array
134
+ edge_features: Array
135
+ batch: Array
136
+
137
+
138
+ # =============================================================================
139
+ # Type Aliases for Operator I/O
140
+ # =============================================================================
141
+
142
+ StateDict = dict[str, Any]
143
+ """State dictionary passed between operator calls."""
144
+
145
+ MetadataDict = dict[str, Any] | None
146
+ """Optional metadata dictionary."""
147
+
148
+ OperatorOutput = tuple[dict[str, Any], StateDict, MetadataDict]
149
+ """Standard operator output: (data, state, metadata)."""
150
+
151
+
152
+ # =============================================================================
153
+ # Protocols for Interfaces
154
+ # =============================================================================
155
+
156
+
157
+ @runtime_checkable
158
+ class DifferentiableOperator(Protocol):
159
+ """Protocol for differentiable operators.
160
+
161
+ All DiffBio operators should implement this interface.
162
+ """
163
+
164
+ def apply(
165
+ self,
166
+ data: dict[str, Any],
167
+ state: StateDict,
168
+ metadata: MetadataDict,
169
+ random_params: Any = None,
170
+ stats: dict[str, Any] | None = None,
171
+ ) -> OperatorOutput:
172
+ """Apply the operator to input data.
173
+
174
+ Args:
175
+ data: Input data dictionary.
176
+ state: Element state.
177
+ metadata: Element metadata.
178
+ random_params: Random parameters for stochastic operations.
179
+ stats: Statistics dictionary.
180
+
181
+ Returns:
182
+ Tuple of (transformed_data, state, metadata).
183
+ """
184
+ ...
185
+
186
+
187
+ @runtime_checkable
188
+ class SequenceEncoder(Protocol):
189
+ """Protocol for sequence encoding/decoding.
190
+
191
+ Implementations should handle conversion between string
192
+ representations and JAX arrays.
193
+ """
194
+
195
+ def encode(self, sequence: str) -> Array:
196
+ """Encode a sequence string to a JAX array.
197
+
198
+ Args:
199
+ sequence: String representation of sequence.
200
+
201
+ Returns:
202
+ Encoded array representation.
203
+ """
204
+ ...
205
+
206
+ def decode(self, encoded: Array) -> str:
207
+ """Decode a JAX array back to a sequence string.
208
+
209
+ Args:
210
+ encoded: Array representation of sequence.
211
+
212
+ Returns:
213
+ String representation.
214
+ """
215
+ ...
216
+
217
+
218
+ @runtime_checkable
219
+ class LossFunction(Protocol):
220
+ """Protocol for loss functions.
221
+
222
+ Loss functions compute scalar losses from predictions and targets.
223
+ """
224
+
225
+ def __call__(
226
+ self,
227
+ predictions: Array,
228
+ targets: Array,
229
+ **kwargs: Any,
230
+ ) -> Float[Array, ""]:
231
+ """Compute the loss.
232
+
233
+ Args:
234
+ predictions: Model predictions.
235
+ targets: Ground truth targets.
236
+ **kwargs: Additional arguments.
237
+
238
+ Returns:
239
+ Scalar loss value.
240
+ """
241
+ ...
242
+
243
+
244
+ @runtime_checkable
245
+ class Regularizer(Protocol):
246
+ """Protocol for regularization functions.
247
+
248
+ Regularizers add penalty terms to loss functions.
249
+ """
250
+
251
+ def __call__(self, params: Any) -> Float[Array, ""]:
252
+ """Compute the regularization penalty.
253
+
254
+ Args:
255
+ params: Parameters to regularize.
256
+
257
+ Returns:
258
+ Scalar regularization term.
259
+ """
260
+ ...