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,240 @@
|
|
|
1
|
+
"""Simplex projection with multiple regularization modes.
|
|
2
|
+
|
|
3
|
+
Projects vectors onto the probability simplex (non-negative, sums to 1)
|
|
4
|
+
using different regularizers that control the smoothness of the
|
|
5
|
+
resulting gradient:
|
|
6
|
+
|
|
7
|
+
- **smooth** (C-infinity): Entropic/softmax regularizer. Closed-form via softmax.
|
|
8
|
+
- **c0** (continuous): Euclidean/L2 regularizer. Solved via threshold algorithm.
|
|
9
|
+
- **c1** (once differentiable): p=3/2 norm regularizer. Closed-form via
|
|
10
|
+
quadratic formula.
|
|
11
|
+
- **c2** (twice differentiable): p=4/3 norm regularizer. Closed-form via
|
|
12
|
+
Cardano's cubic formula.
|
|
13
|
+
|
|
14
|
+
All modes use custom JVP rules for numerically stable gradients.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
from typing import Literal
|
|
20
|
+
|
|
21
|
+
import jax
|
|
22
|
+
import jax.numpy as jnp
|
|
23
|
+
from jax import Array
|
|
24
|
+
|
|
25
|
+
from diffbio.core.soft_ops._utils import canonicalize_axis, validate_softness
|
|
26
|
+
|
|
27
|
+
SimplexMode = Literal["smooth", "c0", "c1", "c2"]
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# --------------------------------------------------------------------------- #
|
|
31
|
+
# C0 projection: Euclidean regularizer (threshold method)
|
|
32
|
+
# --------------------------------------------------------------------------- #
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@jax.custom_jvp
|
|
36
|
+
def _proj_unit_simplex_q2(values: Array) -> Array:
|
|
37
|
+
"""L2 projection onto the unit simplex (1-D, no batch)."""
|
|
38
|
+
n_features = values.shape[0]
|
|
39
|
+
u = jnp.sort(values)[::-1]
|
|
40
|
+
cumsum_u = jnp.cumsum(u)
|
|
41
|
+
ind = jnp.arange(n_features) + 1
|
|
42
|
+
cond = 1.0 / ind + (u - cumsum_u / ind) > 0
|
|
43
|
+
idx = jnp.count_nonzero(cond)
|
|
44
|
+
return jax.nn.relu(1.0 / idx + (values - cumsum_u[idx - 1] / idx))
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
@_proj_unit_simplex_q2.defjvp
|
|
48
|
+
def _proj_unit_simplex_q2_jvp(
|
|
49
|
+
primals: list[Array],
|
|
50
|
+
tangents: list[Array],
|
|
51
|
+
) -> tuple[Array, Array]:
|
|
52
|
+
"""Compute the JVP for L2 simplex projection."""
|
|
53
|
+
(values,) = primals
|
|
54
|
+
(values_dot,) = tangents
|
|
55
|
+
primal_out = _proj_unit_simplex_q2(values)
|
|
56
|
+
supp = primal_out > 0
|
|
57
|
+
card = jnp.count_nonzero(supp)
|
|
58
|
+
tangent_out = supp * values_dot - (jnp.dot(supp, values_dot) / card) * supp
|
|
59
|
+
return primal_out, tangent_out
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
# --------------------------------------------------------------------------- #
|
|
63
|
+
# C1 projection: p=3/2 norm regularizer (quadratic formula)
|
|
64
|
+
# --------------------------------------------------------------------------- #
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _proj_unit_simplex_q3_impl(
|
|
68
|
+
values: Array,
|
|
69
|
+
) -> tuple[Array, Array]:
|
|
70
|
+
"""Closed-form simplex projection for p=3/2 via quadratic formula."""
|
|
71
|
+
n = values.shape[0]
|
|
72
|
+
u = jnp.sort(values)[::-1]
|
|
73
|
+
u0 = u[0]
|
|
74
|
+
u_shift = u - u0
|
|
75
|
+
s_cum = jnp.cumsum(u_shift)
|
|
76
|
+
m2 = jnp.cumsum(u_shift**2)
|
|
77
|
+
k_arr = jnp.arange(1, n + 1, dtype=values.dtype)
|
|
78
|
+
|
|
79
|
+
disc = s_cum**2 - k_arr * (m2 - 1.0)
|
|
80
|
+
theta_k = (s_cum - jnp.sqrt(jnp.maximum(disc, 0.0))) / k_arr
|
|
81
|
+
|
|
82
|
+
cond = u_shift > theta_k
|
|
83
|
+
idx = jnp.count_nonzero(cond)
|
|
84
|
+
theta = theta_k[idx - 1] + u0
|
|
85
|
+
y = jnp.maximum(values - theta, 0.0) ** 2
|
|
86
|
+
return y / jnp.sum(y), theta
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
@jax.custom_jvp
|
|
90
|
+
def _proj_unit_simplex_q3(values: Array) -> Array:
|
|
91
|
+
"""Project onto the unit simplex using p=3/2 norm regularizer."""
|
|
92
|
+
return _proj_unit_simplex_q3_impl(values)[0]
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
@_proj_unit_simplex_q3.defjvp
|
|
96
|
+
def _proj_unit_simplex_q3_jvp(
|
|
97
|
+
primals: list[Array],
|
|
98
|
+
tangents: list[Array],
|
|
99
|
+
) -> tuple[Array, Array]:
|
|
100
|
+
"""Compute the JVP for p=3/2 simplex projection."""
|
|
101
|
+
(values,) = primals
|
|
102
|
+
(values_dot,) = tangents
|
|
103
|
+
primal_out, theta = _proj_unit_simplex_q3_impl(values)
|
|
104
|
+
|
|
105
|
+
supp = (primal_out > 0).astype(values.dtype)
|
|
106
|
+
t = jnp.maximum(values - theta, 0.0)
|
|
107
|
+
w = t * supp
|
|
108
|
+
w_sum = jnp.where(jnp.sum(w) > 0, jnp.sum(w), 1.0)
|
|
109
|
+
|
|
110
|
+
raw_tangent = 2.0 * t * (values_dot - jnp.dot(w, values_dot) / w_sum) * supp
|
|
111
|
+
sum_t2 = jnp.sum(t**2)
|
|
112
|
+
sum_t2 = jnp.where(sum_t2 > 0, sum_t2, 1.0)
|
|
113
|
+
tangent_out = raw_tangent / sum_t2 - primal_out * jnp.sum(raw_tangent) / sum_t2
|
|
114
|
+
return primal_out, tangent_out
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
# --------------------------------------------------------------------------- #
|
|
118
|
+
# C2 projection: p=4/3 norm regularizer (Cardano's cubic formula)
|
|
119
|
+
# --------------------------------------------------------------------------- #
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def _proj_unit_simplex_q4_impl(
|
|
123
|
+
values: Array,
|
|
124
|
+
) -> tuple[Array, Array]:
|
|
125
|
+
"""Closed-form simplex projection for p=4/3 via Cardano's method."""
|
|
126
|
+
n = values.shape[0]
|
|
127
|
+
dtype = values.dtype
|
|
128
|
+
u = jnp.sort(values)[::-1]
|
|
129
|
+
u0 = u[0]
|
|
130
|
+
u_shift = u - u0
|
|
131
|
+
s_cum = jnp.cumsum(u_shift)
|
|
132
|
+
m2 = jnp.cumsum(u_shift**2)
|
|
133
|
+
m3 = jnp.cumsum(u_shift**3)
|
|
134
|
+
k_arr = jnp.arange(1, n + 1, dtype=dtype)
|
|
135
|
+
|
|
136
|
+
c = s_cum / k_arr
|
|
137
|
+
mu2 = m2 - 2.0 * c * s_cum + k_arr * c**2
|
|
138
|
+
mu3 = m3 - 3.0 * c * m2 + 3.0 * c**2 * s_cum - k_arr * c**3
|
|
139
|
+
|
|
140
|
+
p_coeff = 3.0 * mu2 / k_arr
|
|
141
|
+
q_coeff = (1.0 - mu3) / k_arr
|
|
142
|
+
|
|
143
|
+
sp3 = jnp.sqrt(jnp.maximum(p_coeff / 3.0, 0.0))
|
|
144
|
+
denom = 2.0 * jnp.maximum(p_coeff, jnp.finfo(dtype).tiny) * sp3
|
|
145
|
+
big_a = 3.0 * jnp.abs(q_coeff) / denom
|
|
146
|
+
u_hyp = -jnp.sign(q_coeff) * 2.0 * sp3 * jnp.sinh(jnp.arcsinh(big_a) / 3.0)
|
|
147
|
+
u_cbrt = -jnp.sign(q_coeff) * jnp.abs(q_coeff) ** (1.0 / 3.0)
|
|
148
|
+
u_root = jnp.where(
|
|
149
|
+
p_coeff > jnp.finfo(dtype).eps * jnp.maximum(jnp.abs(q_coeff), 1.0),
|
|
150
|
+
u_hyp,
|
|
151
|
+
u_cbrt,
|
|
152
|
+
)
|
|
153
|
+
theta_k = u_root + c
|
|
154
|
+
|
|
155
|
+
cond = u_shift > theta_k
|
|
156
|
+
idx = jnp.count_nonzero(cond)
|
|
157
|
+
theta = theta_k[idx - 1] + u0
|
|
158
|
+
y = jnp.maximum(values - theta, 0.0) ** 3
|
|
159
|
+
return y / jnp.sum(y), theta
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
@jax.custom_jvp
|
|
163
|
+
def _proj_unit_simplex_q4(values: Array) -> Array:
|
|
164
|
+
"""Project onto the unit simplex using p=4/3 norm regularizer."""
|
|
165
|
+
return _proj_unit_simplex_q4_impl(values)[0]
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
@_proj_unit_simplex_q4.defjvp
|
|
169
|
+
def _proj_unit_simplex_q4_jvp(
|
|
170
|
+
primals: list[Array],
|
|
171
|
+
tangents: list[Array],
|
|
172
|
+
) -> tuple[Array, Array]:
|
|
173
|
+
"""Compute the JVP for p=4/3 simplex projection."""
|
|
174
|
+
(values,) = primals
|
|
175
|
+
(values_dot,) = tangents
|
|
176
|
+
primal_out, theta = _proj_unit_simplex_q4_impl(values)
|
|
177
|
+
|
|
178
|
+
supp = (primal_out > 0).astype(values.dtype)
|
|
179
|
+
t = jnp.maximum(values - theta, 0.0)
|
|
180
|
+
w = t**2 * supp
|
|
181
|
+
w_sum = jnp.where(jnp.sum(w) > 0, jnp.sum(w), 1.0)
|
|
182
|
+
|
|
183
|
+
raw_tangent = 3.0 * t**2 * (values_dot - jnp.dot(w, values_dot) / w_sum) * supp
|
|
184
|
+
sum_t3 = jnp.sum(t**3)
|
|
185
|
+
sum_t3 = jnp.where(sum_t3 > 0, sum_t3, 1.0)
|
|
186
|
+
tangent_out = raw_tangent / sum_t3 - primal_out * jnp.sum(raw_tangent) / sum_t3
|
|
187
|
+
return primal_out, tangent_out
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
# --------------------------------------------------------------------------- #
|
|
191
|
+
# Public dispatch function
|
|
192
|
+
# --------------------------------------------------------------------------- #
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
def proj_simplex(
|
|
196
|
+
x: Array,
|
|
197
|
+
axis: int,
|
|
198
|
+
softness: float | Array = 0.1,
|
|
199
|
+
mode: SimplexMode = "smooth",
|
|
200
|
+
) -> Array:
|
|
201
|
+
"""Project ``x`` onto the unit simplex along ``axis``.
|
|
202
|
+
|
|
203
|
+
Solves: ``argmin_y <x, y> + softness * R(y)``
|
|
204
|
+
subject to ``y >= 0, sum(y) = 1``, where ``R(y)`` is determined
|
|
205
|
+
by ``mode``.
|
|
206
|
+
|
|
207
|
+
Args:
|
|
208
|
+
x: Input array of shape ``(..., n, ...)``.
|
|
209
|
+
axis: Axis containing the simplex dimension.
|
|
210
|
+
softness: Regularization strength (> 0). Lower = sharper.
|
|
211
|
+
mode: Regularizer type controlling smoothness:
|
|
212
|
+
``"smooth"`` (C-inf), ``"c0"`` (continuous),
|
|
213
|
+
``"c1"`` (once differentiable), ``"c2"`` (twice differentiable).
|
|
214
|
+
|
|
215
|
+
Returns:
|
|
216
|
+
Projected array on the probability simplex along ``axis``.
|
|
217
|
+
"""
|
|
218
|
+
validate_softness(softness)
|
|
219
|
+
axis = canonicalize_axis(axis, x.ndim)
|
|
220
|
+
scaled = x / softness
|
|
221
|
+
|
|
222
|
+
if mode == "smooth":
|
|
223
|
+
return jax.nn.softmax(scaled, axis=axis)
|
|
224
|
+
|
|
225
|
+
if mode == "c0":
|
|
226
|
+
proj_fn = _proj_unit_simplex_q2
|
|
227
|
+
elif mode == "c1":
|
|
228
|
+
proj_fn = _proj_unit_simplex_q3
|
|
229
|
+
elif mode == "c2":
|
|
230
|
+
proj_fn = _proj_unit_simplex_q4
|
|
231
|
+
else:
|
|
232
|
+
msg = f"Invalid mode: {mode!r}. Must be 'smooth', 'c0', 'c1', or 'c2'."
|
|
233
|
+
raise ValueError(msg)
|
|
234
|
+
|
|
235
|
+
scaled = jnp.moveaxis(scaled, axis, -1)
|
|
236
|
+
*batch_sizes, n = scaled.shape
|
|
237
|
+
scaled = scaled.reshape(-1, n)
|
|
238
|
+
result = jax.vmap(proj_fn)(scaled)
|
|
239
|
+
result = result.reshape(*batch_sizes, n)
|
|
240
|
+
return jnp.moveaxis(result, -1, axis)
|