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,157 @@
1
+ """Neural network utilities for DiffBio.
2
+
3
+ This module provides shared utility functions for building and initializing
4
+ neural network components, ensuring consistency across operators.
5
+ """
6
+
7
+ from typing import TypedDict
8
+
9
+ import jax
10
+ import jax.numpy as jnp
11
+ from flax import nnx
12
+ from jaxtyping import Array
13
+
14
+
15
+ class ArtifexMLPKwargs(TypedDict):
16
+ """Typed shared kwargs for direct Artifex MLP construction."""
17
+
18
+ activation: str
19
+ output_activation: str | None
20
+ use_batch_norm: bool
21
+
22
+
23
+ ARTIFEX_RELU_MLP_KWARGS: ArtifexMLPKwargs = {
24
+ "activation": "relu",
25
+ "output_activation": "relu",
26
+ "use_batch_norm": False,
27
+ }
28
+ ARTIFEX_RELU_BATCH_NORM_MLP_KWARGS: ArtifexMLPKwargs = {
29
+ "activation": "relu",
30
+ "output_activation": "relu",
31
+ "use_batch_norm": True,
32
+ }
33
+ ARTIFEX_GELU_MLP_KWARGS: ArtifexMLPKwargs = {
34
+ "activation": "gelu",
35
+ "output_activation": "gelu",
36
+ "use_batch_norm": False,
37
+ }
38
+ ARTIFEX_GELU_NO_OUTPUT_MLP_KWARGS: ArtifexMLPKwargs = {
39
+ "activation": "gelu",
40
+ "output_activation": None,
41
+ "use_batch_norm": False,
42
+ }
43
+
44
+
45
+ def init_learnable_param(value: float) -> nnx.Param:
46
+ """Initialize a learnable parameter from a scalar value.
47
+
48
+ Args:
49
+ value: Initial scalar value for the parameter.
50
+
51
+ Returns:
52
+ An nnx.Param wrapping a JAX array containing the value.
53
+
54
+ Example:
55
+ ```python
56
+ temperature = init_learnable_param(1.0)
57
+ threshold = init_learnable_param(20.0)
58
+ ```
59
+ """
60
+ return nnx.Param(jnp.array(value))
61
+
62
+
63
+ def ensure_rngs(rngs: nnx.Rngs | None, seed: int = 0) -> nnx.Rngs:
64
+ """Ensure rngs is initialized, creating a default if None.
65
+
66
+ Args:
67
+ rngs: Optional Flax NNX random number generators.
68
+ seed: Seed to use if creating new rngs (default: 0).
69
+
70
+ Returns:
71
+ The provided rngs if not None, otherwise a new nnx.Rngs instance.
72
+
73
+ Example:
74
+ ```python
75
+ rngs = ensure_rngs(rngs) # Use passed rngs or create default
76
+ layer = nnx.Linear(10, 20, rngs=rngs)
77
+ ```
78
+ """
79
+ if rngs is not None:
80
+ return rngs
81
+ return nnx.Rngs(seed)
82
+
83
+
84
+ def get_rng_key(
85
+ rngs: nnx.Rngs | None,
86
+ stream_name: str = "params",
87
+ fallback_seed: int = 0,
88
+ ) -> jax.Array:
89
+ """Get an RNG key from rngs with fallback.
90
+
91
+ Args:
92
+ rngs: Optional Flax NNX random number generators.
93
+ stream_name: Name of the RNG stream to use.
94
+ fallback_seed: Seed to use if rngs is None.
95
+
96
+ Returns:
97
+ A JAX PRNG key.
98
+
99
+ Example:
100
+ ```python
101
+ key = get_rng_key(rngs, "sample")
102
+ noise = jax.random.normal(key, shape)
103
+ ```
104
+ """
105
+ if rngs is not None and stream_name in rngs:
106
+ return getattr(rngs, stream_name)()
107
+ return jax.random.key(fallback_seed)
108
+
109
+
110
+ def extract_windows_1d(
111
+ signal: Array,
112
+ window_size: int,
113
+ pad_mode: str = "edge",
114
+ ) -> Array:
115
+ """Extract sliding windows from a 1D signal with padding.
116
+
117
+ This utility function pads the input signal and extracts overlapping
118
+ windows of the specified size, one centered at each position.
119
+
120
+ Args:
121
+ signal: Input signal of shape (length, features).
122
+ window_size: Size of each window (should be odd for symmetric padding).
123
+ pad_mode: Padding mode for boundaries ("edge", "constant", etc.).
124
+
125
+ Returns:
126
+ Windows of shape (length, window_size, features).
127
+
128
+ Example:
129
+ ```python
130
+ signal = jnp.ones((100, 4)) # 100 positions, 4 features
131
+ windows = extract_windows_1d(signal, window_size=11)
132
+ assert windows.shape == (100, 11, 4)
133
+ ```
134
+ """
135
+ length = signal.shape[0]
136
+ num_features = signal.shape[1]
137
+ half_window = window_size // 2
138
+
139
+ # Pad signal for boundary positions
140
+ padded_signal = jnp.pad(
141
+ signal,
142
+ ((half_window, half_window), (0, 0)),
143
+ mode=pad_mode,
144
+ )
145
+
146
+ # Extract all windows using vmap
147
+ def extract_single_window(pos: Array | int) -> Array:
148
+ return jax.lax.dynamic_slice(
149
+ padded_signal,
150
+ (pos, 0),
151
+ (window_size, num_features),
152
+ )
153
+
154
+ positions = jnp.arange(length)
155
+ all_windows = jax.vmap(extract_single_window)(positions)
156
+
157
+ return all_windows
@@ -0,0 +1,45 @@
1
+ """Quality filtering utilities for DiffBio pipelines.
2
+
3
+ This module provides shared quality filtering functions used across
4
+ multiple pipeline implementations, avoiding code duplication.
5
+ """
6
+
7
+ from datarax.core.operator import OperatorModule
8
+ from jaxtyping import Array, Float
9
+
10
+
11
+ def apply_quality_filter(
12
+ quality_filter: OperatorModule,
13
+ reads: Float[Array, "num_reads read_length 4"],
14
+ quality: Float[Array, "num_reads read_length"],
15
+ ) -> tuple[Float[Array, "num_reads read_length 4"], Float[Array, "num_reads read_length"]]:
16
+ """Apply quality filtering to reads using a differentiable quality filter.
17
+
18
+ Flattens reads and quality scores for per-base filtering, then reshapes
19
+ back to the original dimensions.
20
+
21
+ Args:
22
+ quality_filter: A differentiable quality filter operator
23
+ (e.g., DifferentiableQualityFilter).
24
+ reads: One-hot encoded reads of shape (num_reads, read_length, 4).
25
+ quality: Base quality scores of shape (num_reads, read_length).
26
+
27
+ Returns:
28
+ Tuple of (filtered_reads, filtered_quality) with the same shapes
29
+ as the inputs, where low-quality bases have been soft-masked.
30
+ """
31
+ num_reads, read_length, _ = reads.shape
32
+
33
+ # Flatten for quality filter (treats each base independently)
34
+ reads_flat = reads.reshape(-1, 4)
35
+ quality_flat = quality.reshape(-1)
36
+
37
+ # Apply filter
38
+ filter_data = {"sequence": reads_flat, "quality_scores": quality_flat}
39
+ filtered_result, _, _ = quality_filter.apply(filter_data, {}, None)
40
+
41
+ # Reshape back
42
+ filtered_reads = filtered_result["sequence"].reshape(num_reads, read_length, 4)
43
+ filtered_quality = filtered_result["quality_scores"].reshape(num_reads, read_length)
44
+
45
+ return filtered_reads, filtered_quality