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,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
|
diffbio/utils/quality.py
ADDED
|
@@ -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
|