efax 2.2.0__tar.gz → 2.2.4__tar.gz
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.
- {efax-2.2.0 → efax-2.2.4}/PKG-INFO +36 -11
- {efax-2.2.0 → efax-2.2.4}/README.rst +34 -9
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/categorical.py +2 -1
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/gamma.py +7 -4
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -1
- {efax-2.2.0 → efax-2.2.4}/efax/_src/natural_parametrization.py +25 -6
- {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/support.py +6 -1
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/assembler.py +7 -3
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/estimator.py +6 -2
- {efax-2.2.0 → efax-2.2.4}/efax/_src/tools.py +7 -10
- {efax-2.2.0 → efax-2.2.4}/efax/_src/transform/joint.py +13 -3
- {efax-2.2.0 → efax-2.2.4}/pyproject.toml +3 -4
- {efax-2.2.0 → efax-2.2.4}/efax/__init__.py +3 -3
- {efax-2.2.0 → efax-2.2.4}/efax/_src/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/bernoulli.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/beta.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/chi.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/chi_square.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_von_mises.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/dirichlet.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/dirichlet_common.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/exponential.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/gen_dirichlet.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/geometric.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/inverse_gamma.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/inverse_gaussian.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/log_normal.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/logarithmic.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/fixed_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/negative_binomial.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/negative_binomial_common.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/normal.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/poisson.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/rayleigh.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/von_mises.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/weibull.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/wishart.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/expectation_parametrization.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/conjugate_prior.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/multidimensional.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/samplable.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/iteration.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/has_entropy.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/transformed_parametrization.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/parameter.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/ring.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/parametrization.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/flattener.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/parameter_names.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/parameter_supports.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/transform/__init__.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/_src/types.py +0 -0
- {efax-2.2.0 → efax-2.2.4}/efax/py.typed +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: efax
|
|
3
|
-
Version: 2.2.
|
|
3
|
+
Version: 2.2.4
|
|
4
4
|
Summary: Exponential families for JAX
|
|
5
5
|
Author: Neil Girdhar
|
|
6
6
|
Author-email: Neil Girdhar <mistersheik@gmail.com>
|
|
@@ -23,7 +23,7 @@ Requires-Dist: numpy>=2.0
|
|
|
23
23
|
Requires-Dist: opt-einsum>=3.4
|
|
24
24
|
Requires-Dist: optimistix>=0.0.9
|
|
25
25
|
Requires-Dist: optype[numpy]>=0.8.0
|
|
26
|
-
Requires-Dist: tjax>=1.4.
|
|
26
|
+
Requires-Dist: tjax>=1.4.6
|
|
27
27
|
Requires-Dist: typing-extensions>=4.8
|
|
28
28
|
Maintainer: Neil Girdhar
|
|
29
29
|
Maintainer-email: Neil Girdhar <mistersheik@gmail.com>
|
|
@@ -50,6 +50,12 @@ Description-Content-Type: text/x-rst
|
|
|
50
50
|
:target: https://scientific-python.org/specs/spec-0000/
|
|
51
51
|
:alt: SPEC-0
|
|
52
52
|
:align: center
|
|
53
|
+
.. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json
|
|
54
|
+
:alt: Ruff
|
|
55
|
+
:target: https://github.com/astral-sh/ruff
|
|
56
|
+
.. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ty/main/assets/badge/v0.json
|
|
57
|
+
:alt: ty
|
|
58
|
+
:target: https://github.com/astral-sh/ty
|
|
53
59
|
.. image:: https://img.shields.io/pypi/pyversions/efax
|
|
54
60
|
:alt: PyPI - Python Version
|
|
55
61
|
:align: center
|
|
@@ -161,7 +167,10 @@ Every `NaturalParametrization` has methods:
|
|
|
161
167
|
- `jeffreys_prior_density`, which returns the square root of the Fisher information
|
|
162
168
|
determinant,
|
|
163
169
|
- `characteristic_function`, which evaluates the characteristic function of the sufficient
|
|
164
|
-
statistics via analytic continuation of the log-normalizer
|
|
170
|
+
statistics via analytic continuation of the log-normalizer (note: for distributions such
|
|
171
|
+
as `GammaNP` that keep one or more natural parameters real in `_complexify`, the result
|
|
172
|
+
is approximate — check `NaturalParametrization.characteristic_function_exact` on the
|
|
173
|
+
class before relying on every component), and
|
|
165
174
|
- `kl_divergence`, which is the KL divergence.
|
|
166
175
|
|
|
167
176
|
Every `ExpectationParametrization` has methods:
|
|
@@ -169,6 +178,11 @@ Every `ExpectationParametrization` has methods:
|
|
|
169
178
|
- `to_nat` to convert itself to natural parameters, and
|
|
170
179
|
- `kl_divergence`, which is the KL divergence.
|
|
171
180
|
|
|
181
|
+
The module-level function `expectation_parameters_from_characteristic_function(t, cf_values)`
|
|
182
|
+
estimates the expectation parameters from characteristic-function evaluations via ordinary
|
|
183
|
+
least squares. It is exact for the normal distribution (whose log-CF is linear in the
|
|
184
|
+
frequencies) and a first-order approximation for other families.
|
|
185
|
+
|
|
172
186
|
Some parametrizations inherit from these interfaces:
|
|
173
187
|
|
|
174
188
|
- `HasConjugatePrior` can produce and recover the conjugate prior,
|
|
@@ -313,7 +327,7 @@ EFAX supports the following distributions:
|
|
|
313
327
|
- Dirichlet
|
|
314
328
|
- generalized Dirichlet
|
|
315
329
|
|
|
316
|
-
- on
|
|
330
|
+
- on spheres:
|
|
317
331
|
|
|
318
332
|
- von Mises-Fisher
|
|
319
333
|
- complex von Mises
|
|
@@ -448,17 +462,23 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
448
462
|
return x - 1e-4 * x_bar
|
|
449
463
|
|
|
450
464
|
|
|
451
|
-
|
|
452
|
-
|
|
453
|
-
|
|
465
|
+
# Carry (q, q_bar) as the loop state so the gradient is computed exactly
|
|
466
|
+
# once per step rather than once in cond_fun and once in body_fun.
|
|
467
|
+
type State = tuple[BernoulliNP, BernoulliNP]
|
|
454
468
|
|
|
455
469
|
|
|
456
|
-
def cond_fun(
|
|
457
|
-
q_bar =
|
|
470
|
+
def cond_fun(state: State) -> JaxBooleanArray:
|
|
471
|
+
_, q_bar = state
|
|
458
472
|
total = jnp.sum(parameter_dot_product(q_bar, q_bar))
|
|
459
473
|
return total > 1e-6 # noqa: PLR2004
|
|
460
474
|
|
|
461
475
|
|
|
476
|
+
def body_fun(state: State) -> State:
|
|
477
|
+
q, q_bar = state
|
|
478
|
+
q_new = parameter_map(apply, q, q_bar)
|
|
479
|
+
return q_new, gradient_cross_entropy(target_distribution, q_new)
|
|
480
|
+
|
|
481
|
+
|
|
462
482
|
# The target_distribution is represented as the expectation parameters of a
|
|
463
483
|
# Bernoulli distribution corresponding to probabilities 0.3, 0.4, and 0.7.
|
|
464
484
|
target_distribution = BernoulliEP(jnp.asarray([0.3, 0.4, 0.7]))
|
|
@@ -467,10 +487,13 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
467
487
|
# of a Bernoulli distribution corresponding to log-odds 0, which is probability
|
|
468
488
|
# 0.5.
|
|
469
489
|
initial_predictive_distribution = BernoulliNP(jnp.zeros(3))
|
|
490
|
+
initial_gradient = gradient_cross_entropy(target_distribution,
|
|
491
|
+
initial_predictive_distribution)
|
|
470
492
|
|
|
471
493
|
# Optimize the predictive distribution iteratively.
|
|
472
|
-
predictive_distribution = lax.while_loop(
|
|
473
|
-
|
|
494
|
+
predictive_distribution, _ = lax.while_loop(
|
|
495
|
+
cond_fun, body_fun, (initial_predictive_distribution, initial_gradient)
|
|
496
|
+
)
|
|
474
497
|
|
|
475
498
|
# Compare the optimized predictive distribution with the target value in the
|
|
476
499
|
# same natural parametrization.
|
|
@@ -535,6 +558,8 @@ instead.
|
|
|
535
558
|
# Dirichlet distribution.
|
|
536
559
|
|
|
537
560
|
# Take the mean over the first axis.
|
|
561
|
+
# parameter_mean averages only variable parameters; fixed parameters (e.g.
|
|
562
|
+
# the failure count in NegativeBinomialNP) are preserved unchanged.
|
|
538
563
|
ss_mean = parameter_mean(ss, axis=0) # ss_mean also has type DirichletEP.
|
|
539
564
|
|
|
540
565
|
# Convert this back to the natural parametrization.
|
|
@@ -17,6 +17,12 @@
|
|
|
17
17
|
:target: https://scientific-python.org/specs/spec-0000/
|
|
18
18
|
:alt: SPEC-0
|
|
19
19
|
:align: center
|
|
20
|
+
.. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json
|
|
21
|
+
:alt: Ruff
|
|
22
|
+
:target: https://github.com/astral-sh/ruff
|
|
23
|
+
.. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ty/main/assets/badge/v0.json
|
|
24
|
+
:alt: ty
|
|
25
|
+
:target: https://github.com/astral-sh/ty
|
|
20
26
|
.. image:: https://img.shields.io/pypi/pyversions/efax
|
|
21
27
|
:alt: PyPI - Python Version
|
|
22
28
|
:align: center
|
|
@@ -128,7 +134,10 @@ Every `NaturalParametrization` has methods:
|
|
|
128
134
|
- `jeffreys_prior_density`, which returns the square root of the Fisher information
|
|
129
135
|
determinant,
|
|
130
136
|
- `characteristic_function`, which evaluates the characteristic function of the sufficient
|
|
131
|
-
statistics via analytic continuation of the log-normalizer
|
|
137
|
+
statistics via analytic continuation of the log-normalizer (note: for distributions such
|
|
138
|
+
as `GammaNP` that keep one or more natural parameters real in `_complexify`, the result
|
|
139
|
+
is approximate — check `NaturalParametrization.characteristic_function_exact` on the
|
|
140
|
+
class before relying on every component), and
|
|
132
141
|
- `kl_divergence`, which is the KL divergence.
|
|
133
142
|
|
|
134
143
|
Every `ExpectationParametrization` has methods:
|
|
@@ -136,6 +145,11 @@ Every `ExpectationParametrization` has methods:
|
|
|
136
145
|
- `to_nat` to convert itself to natural parameters, and
|
|
137
146
|
- `kl_divergence`, which is the KL divergence.
|
|
138
147
|
|
|
148
|
+
The module-level function `expectation_parameters_from_characteristic_function(t, cf_values)`
|
|
149
|
+
estimates the expectation parameters from characteristic-function evaluations via ordinary
|
|
150
|
+
least squares. It is exact for the normal distribution (whose log-CF is linear in the
|
|
151
|
+
frequencies) and a first-order approximation for other families.
|
|
152
|
+
|
|
139
153
|
Some parametrizations inherit from these interfaces:
|
|
140
154
|
|
|
141
155
|
- `HasConjugatePrior` can produce and recover the conjugate prior,
|
|
@@ -280,7 +294,7 @@ EFAX supports the following distributions:
|
|
|
280
294
|
- Dirichlet
|
|
281
295
|
- generalized Dirichlet
|
|
282
296
|
|
|
283
|
-
- on
|
|
297
|
+
- on spheres:
|
|
284
298
|
|
|
285
299
|
- von Mises-Fisher
|
|
286
300
|
- complex von Mises
|
|
@@ -415,17 +429,23 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
415
429
|
return x - 1e-4 * x_bar
|
|
416
430
|
|
|
417
431
|
|
|
418
|
-
|
|
419
|
-
|
|
420
|
-
|
|
432
|
+
# Carry (q, q_bar) as the loop state so the gradient is computed exactly
|
|
433
|
+
# once per step rather than once in cond_fun and once in body_fun.
|
|
434
|
+
type State = tuple[BernoulliNP, BernoulliNP]
|
|
421
435
|
|
|
422
436
|
|
|
423
|
-
def cond_fun(
|
|
424
|
-
q_bar =
|
|
437
|
+
def cond_fun(state: State) -> JaxBooleanArray:
|
|
438
|
+
_, q_bar = state
|
|
425
439
|
total = jnp.sum(parameter_dot_product(q_bar, q_bar))
|
|
426
440
|
return total > 1e-6 # noqa: PLR2004
|
|
427
441
|
|
|
428
442
|
|
|
443
|
+
def body_fun(state: State) -> State:
|
|
444
|
+
q, q_bar = state
|
|
445
|
+
q_new = parameter_map(apply, q, q_bar)
|
|
446
|
+
return q_new, gradient_cross_entropy(target_distribution, q_new)
|
|
447
|
+
|
|
448
|
+
|
|
429
449
|
# The target_distribution is represented as the expectation parameters of a
|
|
430
450
|
# Bernoulli distribution corresponding to probabilities 0.3, 0.4, and 0.7.
|
|
431
451
|
target_distribution = BernoulliEP(jnp.asarray([0.3, 0.4, 0.7]))
|
|
@@ -434,10 +454,13 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
434
454
|
# of a Bernoulli distribution corresponding to log-odds 0, which is probability
|
|
435
455
|
# 0.5.
|
|
436
456
|
initial_predictive_distribution = BernoulliNP(jnp.zeros(3))
|
|
457
|
+
initial_gradient = gradient_cross_entropy(target_distribution,
|
|
458
|
+
initial_predictive_distribution)
|
|
437
459
|
|
|
438
460
|
# Optimize the predictive distribution iteratively.
|
|
439
|
-
predictive_distribution = lax.while_loop(
|
|
440
|
-
|
|
461
|
+
predictive_distribution, _ = lax.while_loop(
|
|
462
|
+
cond_fun, body_fun, (initial_predictive_distribution, initial_gradient)
|
|
463
|
+
)
|
|
441
464
|
|
|
442
465
|
# Compare the optimized predictive distribution with the target value in the
|
|
443
466
|
# same natural parametrization.
|
|
@@ -502,6 +525,8 @@ instead.
|
|
|
502
525
|
# Dirichlet distribution.
|
|
503
526
|
|
|
504
527
|
# Take the mean over the first axis.
|
|
528
|
+
# parameter_mean averages only variable parameters; fixed parameters (e.g.
|
|
529
|
+
# the failure count in NegativeBinomialNP) are preserved unchanged.
|
|
505
530
|
ss_mean = parameter_mean(ss, axis=0) # ss_mean also has type DirichletEP.
|
|
506
531
|
|
|
507
532
|
# Convert this back to the natural parametrization.
|
|
@@ -108,7 +108,8 @@ class CategoricalEP(
|
|
|
108
108
|
"""The expectation parametrization of the categorical distribution.
|
|
109
109
|
|
|
110
110
|
Args:
|
|
111
|
-
probability: The probability vector with the final element omitted, i.e.,
|
|
111
|
+
probability: The probability vector with the final element omitted, i.e.,
|
|
112
|
+
[p_i]_{i in 1...n-1}.
|
|
112
113
|
"""
|
|
113
114
|
|
|
114
115
|
probability: JaxRealArray = distribution_parameter(
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
-
from typing import Self, override
|
|
3
|
+
from typing import ClassVar, Self, override
|
|
4
4
|
|
|
5
5
|
import jax.random as jr
|
|
6
6
|
import jax.scipy.special as jss
|
|
@@ -39,6 +39,8 @@ class GammaNP(
|
|
|
39
39
|
shape_minus_one: The shape minus one.
|
|
40
40
|
"""
|
|
41
41
|
|
|
42
|
+
characteristic_function_exact: ClassVar[bool] = False
|
|
43
|
+
|
|
42
44
|
negative_rate: JaxRealArray = distribution_parameter(ScalarSupport(ring=negative_support))
|
|
43
45
|
shape_minus_one: JaxRealArray = distribution_parameter(
|
|
44
46
|
ScalarSupport(ring=RealField(minimum=-1.0))
|
|
@@ -83,9 +85,10 @@ class GammaNP(
|
|
|
83
85
|
def _complexify(self, t: Self) -> Self: # ty: ignore
|
|
84
86
|
"""Only complexify negative_rate; jss.gammaln does not accept complex inputs.
|
|
85
87
|
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
88
|
+
shape_minus_one is kept real, so ``characteristic_function`` cannot
|
|
89
|
+
query the log-x sufficient statistic: that component returns 1 instead
|
|
90
|
+
of the true E[exp(i·t·log x)]. Consequently,
|
|
91
|
+
``characteristic_function_exact = False`` on this class.
|
|
89
92
|
"""
|
|
90
93
|
return type(self)(self.negative_rate + 1j * t.negative_rate, self.shape_minus_one)
|
|
91
94
|
|
|
@@ -55,7 +55,8 @@ class OptimistixRootFinder(ExpToNatMinimizer):
|
|
|
55
55
|
|
|
56
56
|
|
|
57
57
|
default_minimizer = OptimistixRootFinder(
|
|
58
|
-
solver=optx.Newton[JaxRealArray, JaxRealArray, None](rtol=0.0, atol=1e-7),
|
|
58
|
+
solver=optx.Newton[JaxRealArray, JaxRealArray, None](rtol=0.0, atol=1e-7), # ty: ignore
|
|
59
|
+
max_steps=1000,
|
|
59
60
|
)
|
|
60
61
|
|
|
61
62
|
|
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
3
|
from abc import abstractmethod
|
|
4
|
-
from typing import TYPE_CHECKING, Any, Generic, Self, final, get_type_hints
|
|
4
|
+
from typing import TYPE_CHECKING, Any, ClassVar, Generic, Self, final, get_type_hints
|
|
5
5
|
|
|
6
6
|
import jax
|
|
7
7
|
import jax.numpy as jnp
|
|
@@ -53,8 +53,18 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
|
|
|
53
53
|
|
|
54
54
|
The motivation for the natural parametrization is combining and scaling independent predictive
|
|
55
55
|
evidence. In the natural parametrization, these operations correspond to scaling and addition.
|
|
56
|
+
|
|
57
|
+
Class variables:
|
|
58
|
+
characteristic_function_exact: True when `_complexify` shifts every natural parameter
|
|
59
|
+
into the complex plane, so `characteristic_function` queries all sufficient
|
|
60
|
+
statistics exactly. Set to False in subclasses that must keep one or more
|
|
61
|
+
parameters real (e.g. because `log_normalizer` calls a function such as
|
|
62
|
+
``gammaln`` that does not accept complex inputs); those components return 1
|
|
63
|
+
instead of the true characteristic-function value.
|
|
56
64
|
"""
|
|
57
65
|
|
|
66
|
+
characteristic_function_exact: ClassVar[bool] = True
|
|
67
|
+
|
|
58
68
|
@abstract_custom_jvp(_log_normalizer_jvp)
|
|
59
69
|
@abstract_jit
|
|
60
70
|
@abstractmethod
|
|
@@ -196,11 +206,13 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
|
|
|
196
206
|
def _complexify(self, t: Self) -> Self:
|
|
197
207
|
"""Shift natural parameters into the complex plane: η → η + i·t.
|
|
198
208
|
|
|
199
|
-
The default shifts every field
|
|
200
|
-
``log_normalizer``
|
|
201
|
-
(e.g. ``
|
|
202
|
-
``
|
|
203
|
-
|
|
209
|
+
The default shifts every field, which is correct for any distribution
|
|
210
|
+
whose ``log_normalizer`` accepts fully complex inputs. Subclasses that
|
|
211
|
+
must keep one or more parameters real (e.g. because ``log_normalizer``
|
|
212
|
+
calls ``gammaln``, which rejects complex inputs) must override this
|
|
213
|
+
method **and** set ``characteristic_function_exact = False`` on the
|
|
214
|
+
class. For those parameters, ``characteristic_function`` returns 1
|
|
215
|
+
instead of the true E[exp(i·⟨t, T(x)⟩)] component.
|
|
204
216
|
|
|
205
217
|
Args:
|
|
206
218
|
t: Imaginary displacements, same pytree structure as self.
|
|
@@ -237,6 +249,13 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
|
|
|
237
249
|
t_grid = PoissonNP(omegas) # shape (k,)
|
|
238
250
|
phi = p.characteristic_function(t_grid) # shape (k,)
|
|
239
251
|
|
|
252
|
+
Note: when ``type(self).characteristic_function_exact`` is False, one
|
|
253
|
+
or more natural parameters are kept real in ``_complexify`` because
|
|
254
|
+
``log_normalizer`` uses a function that rejects complex inputs. The
|
|
255
|
+
corresponding sufficient-statistic components return 1 rather than
|
|
256
|
+
the true value. Check the class attribute before relying on the
|
|
257
|
+
result for those components.
|
|
258
|
+
|
|
240
259
|
Args:
|
|
241
260
|
t: Frequencies in natural-parameter space, same structure as self.
|
|
242
261
|
|
|
@@ -31,7 +31,6 @@ def triangular_number_index(k: int, /) -> int | None:
|
|
|
31
31
|
|
|
32
32
|
|
|
33
33
|
class Support:
|
|
34
|
-
@override
|
|
35
34
|
def __init__(self, *, ring: Ring = real_field) -> None:
|
|
36
35
|
super().__init__()
|
|
37
36
|
self.ring = ring
|
|
@@ -280,6 +279,12 @@ class SquareMatrixSupport(Support):
|
|
|
280
279
|
x = self.ring.unflattened(y, map_from_plane=map_from_plane)
|
|
281
280
|
return xp.reshape(x, x.shape[:-1] + self.shape(dimensions))
|
|
282
281
|
|
|
282
|
+
@override
|
|
283
|
+
def generate(
|
|
284
|
+
self, xp: Namespace, rng: Generator, shape: Shape, safety: float, dimensions: int
|
|
285
|
+
) -> JaxRealArray:
|
|
286
|
+
return self.ring.generate(xp, rng, (*shape, dimensions, dimensions), safety)
|
|
287
|
+
|
|
283
288
|
|
|
284
289
|
class CircularBoundedSupport(VectorSupport):
|
|
285
290
|
def __init__(self, radius: float) -> None:
|
|
@@ -68,7 +68,9 @@ class Assembler(Generic[P]):
|
|
|
68
68
|
|
|
69
69
|
infos = []
|
|
70
70
|
for info in self.infos:
|
|
71
|
-
|
|
71
|
+
if not issubclass(info.type_, ExpectationParametrization):
|
|
72
|
+
msg = f"{info.type_.__name__} is not an EP"
|
|
73
|
+
raise TypeError(msg)
|
|
72
74
|
infos.append(
|
|
73
75
|
SubDistributionInfo(
|
|
74
76
|
info.path,
|
|
@@ -85,7 +87,9 @@ class Assembler(Generic[P]):
|
|
|
85
87
|
|
|
86
88
|
infos: list[SubDistributionInfo] = []
|
|
87
89
|
for info in self.infos:
|
|
88
|
-
|
|
90
|
+
if not issubclass(info.type_, NaturalParametrization):
|
|
91
|
+
msg = f"{info.type_.__name__} is not an NP"
|
|
92
|
+
raise TypeError(msg)
|
|
89
93
|
infos.append(
|
|
90
94
|
SubDistributionInfo(
|
|
91
95
|
info.path,
|
|
@@ -129,7 +133,7 @@ class Assembler(Generic[P]):
|
|
|
129
133
|
]
|
|
130
134
|
q_values = parameters(q).values()
|
|
131
135
|
q_params_as_p = dict(zip(p_paths, q_values, strict=True))
|
|
132
|
-
return self.assemble(q_params_as_p)
|
|
136
|
+
return self.assemble(q_params_as_p) # ty: ignore
|
|
133
137
|
|
|
134
138
|
def domain_support(self) -> dict[Path, Support]:
|
|
135
139
|
"""Return the domain support constraints for each simple sub-distribution in the tree."""
|
|
@@ -45,7 +45,9 @@ class Estimator(Assembler[P]):
|
|
|
45
45
|
ExpectationParametrization,
|
|
46
46
|
)
|
|
47
47
|
|
|
48
|
-
|
|
48
|
+
if not issubclass(type_p, ExpectationParametrization):
|
|
49
|
+
msg = f"{type_p.__name__} is not an EP"
|
|
50
|
+
raise TypeError(msg)
|
|
49
51
|
return Estimator(
|
|
50
52
|
[SubDistributionInfo((), type_p, 0, [])],
|
|
51
53
|
{(name,): value for name, value in fixed_parameters.items()},
|
|
@@ -63,7 +65,9 @@ class Estimator(Assembler[P]):
|
|
|
63
65
|
)
|
|
64
66
|
|
|
65
67
|
infos = cls.create_assembler(p).infos
|
|
66
|
-
|
|
68
|
+
if not isinstance(p, ExpectationParametrization):
|
|
69
|
+
msg = f"{type(p).__name__} is not an EP"
|
|
70
|
+
raise TypeError(msg)
|
|
67
71
|
fixed_parameters = parameters(p, fixed=True)
|
|
68
72
|
return cls(infos, fixed_parameters)
|
|
69
73
|
|
|
@@ -4,7 +4,7 @@ from collections import defaultdict
|
|
|
4
4
|
from collections.abc import Callable, Iterable, Mapping
|
|
5
5
|
from functools import reduce
|
|
6
6
|
from itertools import starmap
|
|
7
|
-
from typing import
|
|
7
|
+
from typing import TypeVar
|
|
8
8
|
|
|
9
9
|
from array_api_compat import array_namespace
|
|
10
10
|
from jax import jit
|
|
@@ -17,7 +17,7 @@ from .types import Axis
|
|
|
17
17
|
|
|
18
18
|
|
|
19
19
|
@jit
|
|
20
|
-
def parameter_dot_product(x:
|
|
20
|
+
def parameter_dot_product(x: Distribution, y: Distribution, /) -> JaxRealArray:
|
|
21
21
|
"""Return the vectorized dot product over all of the variable parameters."""
|
|
22
22
|
|
|
23
23
|
def dotted_fields() -> Iterable[JaxRealArray]:
|
|
@@ -38,12 +38,13 @@ T = TypeVar("T", bound=Distribution)
|
|
|
38
38
|
|
|
39
39
|
|
|
40
40
|
def parameter_mean[T: Distribution](x: T, /, *, axis: Axis | None = None) -> T:
|
|
41
|
-
"""Return the mean of the parameters
|
|
41
|
+
"""Return the mean of the variable parameters; fixed parameters are preserved unchanged."""
|
|
42
42
|
xp = array_namespace(x)
|
|
43
43
|
structure = Assembler.create_assembler(x)
|
|
44
|
-
|
|
45
|
-
|
|
46
|
-
|
|
44
|
+
fixed = parameters(x, fixed=True)
|
|
45
|
+
variable = parameters(x, fixed=False)
|
|
46
|
+
averaged = {path: xp.mean(value, axis=axis) for path, value in variable.items()}
|
|
47
|
+
return structure.assemble({**fixed, **averaged})
|
|
47
48
|
|
|
48
49
|
|
|
49
50
|
def parameter_map[T: Distribution](
|
|
@@ -80,7 +81,3 @@ def _parameter_dot_product(x: JaxComplexArray, y: JaxComplexArray, n_axes: int)
|
|
|
80
81
|
xp = array_namespace(x, y)
|
|
81
82
|
axes = tuple(range(-n_axes, 0))
|
|
82
83
|
return xp.sum(xp.real(x * xp.conj(y)), axis=axes)
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
if TYPE_CHECKING:
|
|
86
|
-
from .natural_parametrization import NaturalParametrization
|
|
@@ -50,9 +50,19 @@ class JointDistribution(Distribution):
|
|
|
50
50
|
@property
|
|
51
51
|
@override
|
|
52
52
|
def shape(self) -> Shape:
|
|
53
|
-
for
|
|
54
|
-
|
|
55
|
-
|
|
53
|
+
shapes = [d.shape for d in self._sub_distributions.values()]
|
|
54
|
+
if not shapes:
|
|
55
|
+
msg = "JointDistribution has no sub-distributions"
|
|
56
|
+
raise ValueError(msg)
|
|
57
|
+
shape = shapes[0]
|
|
58
|
+
if any(s != shape for s in shapes[1:]):
|
|
59
|
+
names = list(self._sub_distributions)
|
|
60
|
+
msg = (
|
|
61
|
+
f"Sub-distribution shapes are inconsistent: "
|
|
62
|
+
f"{dict(zip(names, shapes, strict=True))}"
|
|
63
|
+
)
|
|
64
|
+
raise ValueError(msg)
|
|
65
|
+
return shape
|
|
56
66
|
|
|
57
67
|
|
|
58
68
|
@dataclass
|
|
@@ -6,19 +6,18 @@ build-backend = "uv_build"
|
|
|
6
6
|
dev = [
|
|
7
7
|
"jupyter>=1",
|
|
8
8
|
"lefthook>=2",
|
|
9
|
-
"pyright>=1.1.408",
|
|
10
9
|
"pytest-xdist[psutil]>=3",
|
|
11
10
|
"pytest>=9",
|
|
12
11
|
"ruff>=0.15",
|
|
13
12
|
"scipy-stubs>=1.15",
|
|
14
13
|
"scipy>=1.15",
|
|
15
14
|
"toml-sort>=0.24",
|
|
16
|
-
"ty>=0.0.
|
|
15
|
+
"ty>=0.0.29"
|
|
17
16
|
]
|
|
18
17
|
|
|
19
18
|
[project]
|
|
20
19
|
name = "efax"
|
|
21
|
-
version = "2.2.
|
|
20
|
+
version = "2.2.4"
|
|
22
21
|
description = "Exponential families for JAX"
|
|
23
22
|
readme = "README.rst"
|
|
24
23
|
requires-python = ">=3.12, <3.15"
|
|
@@ -46,7 +45,7 @@ dependencies = [
|
|
|
46
45
|
"opt-einsum>=3.4",
|
|
47
46
|
"optimistix>=0.0.9",
|
|
48
47
|
"optype[numpy]>=0.8.0",
|
|
49
|
-
"tjax>=1.4.
|
|
48
|
+
"tjax>=1.4.6",
|
|
50
49
|
"typing_extensions>=4.8"
|
|
51
50
|
]
|
|
52
51
|
|
|
@@ -2,6 +2,7 @@
|
|
|
2
2
|
|
|
3
3
|
from ._src.distributions.bernoulli import BernoulliEP, BernoulliNP
|
|
4
4
|
from ._src.distributions.beta import BetaEP, BetaNP
|
|
5
|
+
from ._src.distributions.categorical import CategoricalEP, CategoricalNP
|
|
5
6
|
from ._src.distributions.chi import ChiEP, ChiNP
|
|
6
7
|
from ._src.distributions.chi_square import ChiSquareEP, ChiSquareNP
|
|
7
8
|
from ._src.distributions.cmvn.circularly_symmetric import (
|
|
@@ -31,7 +32,6 @@ from ._src.distributions.log_normal.unit_variance import (
|
|
|
31
32
|
UnitVarianceLogNormalNP,
|
|
32
33
|
)
|
|
33
34
|
from ._src.distributions.logarithmic import LogarithmicEP, LogarithmicNP
|
|
34
|
-
from ._src.distributions.categorical import CategoricalEP, CategoricalNP
|
|
35
35
|
from ._src.distributions.multivariate_normal.arbitrary import (
|
|
36
36
|
MultivariateNormalEP,
|
|
37
37
|
MultivariateNormalNP,
|
|
@@ -105,6 +105,8 @@ __all__ = [
|
|
|
105
105
|
"BetaEP",
|
|
106
106
|
"BetaNP",
|
|
107
107
|
"BooleanRing",
|
|
108
|
+
"CategoricalEP",
|
|
109
|
+
"CategoricalNP",
|
|
108
110
|
"ChiEP",
|
|
109
111
|
"ChiNP",
|
|
110
112
|
"ChiSquareEP",
|
|
@@ -156,8 +158,6 @@ __all__ = [
|
|
|
156
158
|
"LogarithmicEP",
|
|
157
159
|
"LogarithmicNP",
|
|
158
160
|
"Multidimensional",
|
|
159
|
-
"CategoricalEP",
|
|
160
|
-
"CategoricalNP",
|
|
161
161
|
"MultivariateDiagonalNormalEP",
|
|
162
162
|
"MultivariateDiagonalNormalNP",
|
|
163
163
|
"MultivariateDiagonalNormalVP",
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|