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