efax 2.1.2__tar.gz → 2.2.3__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.1.2 → efax-2.2.3}/PKG-INFO +156 -59
- {efax-2.1.2 → efax-2.2.3}/README.rst +148 -51
- {efax-2.1.2 → efax-2.2.3}/efax/__init__.py +7 -7
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/bernoulli.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/beta.py +1 -0
- efax-2.1.2/efax/_src/distributions/multinomial.py → efax-2.2.3/efax/_src/distributions/categorical.py +20 -18
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/chi.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/unit_variance.py +1 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/unit_variance.py +1 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/dirichlet.py +1 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/exponential.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/gamma.py +8 -4
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/geometric.py +1 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/inverse_gamma.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/log_normal.py +2 -2
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/unit_variance.py +4 -4
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/fixed_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/unit_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/negative_binomial.py +1 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/unit_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/poisson.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/rayleigh.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/softplus.py +2 -2
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/unit_variance.py +4 -4
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/von_mises.py +3 -3
- {efax-2.1.2 → efax-2.2.3}/efax/_src/expectation_parametrization.py +7 -6
- efax-2.2.3/efax/_src/interfaces/conjugate_prior.py +72 -0
- efax-2.2.3/efax/_src/interfaces/multidimensional.py +18 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/interfaces/samplable.py +9 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +12 -3
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/has_entropy.py +23 -8
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/transformed_parametrization.py +36 -9
- {efax-2.1.2 → efax-2.2.3}/efax/_src/natural_parametrization.py +43 -18
- {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/support.py +7 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/parametrization.py +24 -5
- efax-2.1.2/efax/_src/structure/structure.py → efax-2.2.3/efax/_src/structure/assembler.py +54 -26
- {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/estimator.py +57 -25
- {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/flattener.py +46 -34
- {efax-2.1.2 → efax-2.2.3}/efax/_src/tools.py +10 -13
- {efax-2.1.2 → efax-2.2.3}/efax/_src/transform/joint.py +13 -3
- {efax-2.1.2 → efax-2.2.3}/pyproject.toml +8 -9
- efax-2.1.2/.editorconfig +0 -19
- efax-2.1.2/.gitignore +0 -2
- efax-2.1.2/LICENSE +0 -201
- efax-2.1.2/efax/_src/interfaces/conjugate_prior.py +0 -53
- efax-2.1.2/efax/_src/interfaces/multidimensional.py +0 -11
- efax-2.1.2/examples/.editorconfig +0 -2
- efax-2.1.2/examples/__init__.py +0 -1
- efax-2.1.2/examples/bayesian_evidence_combination.py +0 -48
- efax-2.1.2/examples/cross_entropy.py +0 -33
- efax-2.1.2/examples/maximum_likelihood_estimation.py +0 -47
- efax-2.1.2/examples/optimization.py +0 -79
- efax-2.1.2/exponential_families.pdf +3 -29318
- efax-2.1.2/lefthook.yml +0 -18
- efax-2.1.2/tests/__init__.py +0 -1
- efax-2.1.2/tests/conftest.py +0 -152
- efax-2.1.2/tests/create_info.py +0 -857
- efax-2.1.2/tests/distribution_info.py +0 -116
- efax-2.1.2/tests/match_scipy/__init__.py +0 -1
- efax-2.1.2/tests/match_scipy/test_entropy.py +0 -33
- efax-2.1.2/tests/match_scipy/test_maximum_likelihood_estimation.py +0 -78
- efax-2.1.2/tests/match_scipy/test_pdf.py +0 -71
- efax-2.1.2/tests/scipy_replacement/__init__.py +0 -1
- efax-2.1.2/tests/scipy_replacement/base.py +0 -24
- efax-2.1.2/tests/scipy_replacement/complex_multivariate_normal.py +0 -160
- efax-2.1.2/tests/scipy_replacement/complex_normal.py +0 -157
- efax-2.1.2/tests/scipy_replacement/dirichlet.py +0 -92
- efax-2.1.2/tests/scipy_replacement/joint.py +0 -33
- efax-2.1.2/tests/scipy_replacement/multinomial.py +0 -26
- efax-2.1.2/tests/scipy_replacement/multivariate_normal.py +0 -54
- efax-2.1.2/tests/scipy_replacement/shaped_distribution.py +0 -92
- efax-2.1.2/tests/scipy_replacement/von_mises.py +0 -44
- efax-2.1.2/tests/scipy_replacement/wishart.py +0 -32
- efax-2.1.2/tests/softplus.py +0 -16
- efax-2.1.2/tests/test_characteristic_function.py +0 -138
- efax-2.1.2/tests/test_complex_normal.py +0 -105
- efax-2.1.2/tests/test_complex_von_mises.py +0 -27
- efax-2.1.2/tests/test_conjugate_prior.py +0 -95
- efax-2.1.2/tests/test_conversion.py +0 -42
- efax-2.1.2/tests/test_degenerate.py +0 -38
- efax-2.1.2/tests/test_entropy_gradient.py +0 -90
- efax-2.1.2/tests/test_ep_from_cf.py +0 -127
- efax-2.1.2/tests/test_fisher_information.py +0 -93
- efax-2.1.2/tests/test_flatten.py +0 -42
- efax-2.1.2/tests/test_gradient_log_normalizer.py +0 -111
- efax-2.1.2/tests/test_hessian.py +0 -97
- efax-2.1.2/tests/test_jax_quirks.py +0 -14
- efax-2.1.2/tests/test_kl.py +0 -85
- efax-2.1.2/tests/test_reparametrization_trick.py +0 -22
- efax-2.1.2/tests/test_sampling.py +0 -193
- efax-2.1.2/tests/test_scipy_distributions.py +0 -32
- efax-2.1.2/tests/test_shapes.py +0 -38
- efax-2.1.2/uv.lock +0 -2199
- {efax-2.1.2 → efax-2.2.3}/efax/_src/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/chi_square.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_von_mises.py +1 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/dirichlet_common.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/gen_dirichlet.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/inverse_gaussian.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/logarithmic.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/negative_binomial_common.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/normal.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/weibull.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/wishart.py +1 -1
- {efax-2.1.2 → efax-2.2.3}/efax/_src/interfaces/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/iteration.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/parameter.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/ring.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/parameter_names.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/parameter_supports.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/transform/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/_src/types.py +0 -0
- {efax-2.1.2 → efax-2.2.3}/efax/py.typed +0 -0
|
@@ -1,25 +1,21 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: efax
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.2.3
|
|
4
4
|
Summary: Exponential families for JAX
|
|
5
|
-
|
|
5
|
+
Author: Neil Girdhar
|
|
6
6
|
Author-email: Neil Girdhar <mistersheik@gmail.com>
|
|
7
|
-
Maintainer-email: Neil Girdhar <mistersheik@gmail.com>
|
|
8
7
|
License-Expression: Apache-2.0
|
|
9
|
-
License-File: LICENSE
|
|
10
8
|
Classifier: Development Status :: 5 - Production/Stable
|
|
11
9
|
Classifier: Intended Audience :: Science/Research
|
|
12
|
-
Classifier: License :: OSI Approved :: Apache Software License
|
|
13
10
|
Classifier: Operating System :: OS Independent
|
|
14
|
-
Classifier: Programming Language :: Python
|
|
15
11
|
Classifier: Programming Language :: Python :: 3
|
|
16
12
|
Classifier: Programming Language :: Python :: 3.12
|
|
17
13
|
Classifier: Programming Language :: Python :: 3.13
|
|
18
14
|
Classifier: Programming Language :: Python :: 3.14
|
|
15
|
+
Classifier: Programming Language :: Python
|
|
19
16
|
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
20
17
|
Classifier: Topic :: Software Development :: Libraries :: Python Modules
|
|
21
18
|
Classifier: Typing :: Typed
|
|
22
|
-
Requires-Python: <3.15,>=3.12
|
|
23
19
|
Requires-Dist: array-api-compat>=1.10
|
|
24
20
|
Requires-Dist: array-api-extra>=0.9.2
|
|
25
21
|
Requires-Dist: jax>=0.9.0
|
|
@@ -27,8 +23,12 @@ Requires-Dist: numpy>=2.0
|
|
|
27
23
|
Requires-Dist: opt-einsum>=3.4
|
|
28
24
|
Requires-Dist: optimistix>=0.0.9
|
|
29
25
|
Requires-Dist: optype[numpy]>=0.8.0
|
|
30
|
-
Requires-Dist: tjax>=1.4.
|
|
26
|
+
Requires-Dist: tjax>=1.4.6
|
|
31
27
|
Requires-Dist: typing-extensions>=4.8
|
|
28
|
+
Maintainer: Neil Girdhar
|
|
29
|
+
Maintainer-email: Neil Girdhar <mistersheik@gmail.com>
|
|
30
|
+
Requires-Python: >=3.12, <3.15
|
|
31
|
+
Project-URL: source, https://github.com/NeilGirdhar/efax
|
|
32
32
|
Description-Content-Type: text/x-rst
|
|
33
33
|
|
|
34
34
|
.. role:: bash(code)
|
|
@@ -37,6 +37,8 @@ Description-Content-Type: text/x-rst
|
|
|
37
37
|
.. role:: python(code)
|
|
38
38
|
:language: python
|
|
39
39
|
|
|
40
|
+
.. default-role:: python
|
|
41
|
+
|
|
40
42
|
.. image:: https://img.shields.io/pypi/v/efax
|
|
41
43
|
:target: https://pypi.org/project/efax/
|
|
42
44
|
:alt: PyPI - Version
|
|
@@ -48,6 +50,12 @@ Description-Content-Type: text/x-rst
|
|
|
48
50
|
:target: https://scientific-python.org/specs/spec-0000/
|
|
49
51
|
:alt: SPEC-0
|
|
50
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
|
|
51
59
|
.. image:: https://img.shields.io/pypi/pyversions/efax
|
|
52
60
|
:alt: PyPI - Python Version
|
|
53
61
|
:align: center
|
|
@@ -81,7 +89,7 @@ Framework
|
|
|
81
89
|
=========
|
|
82
90
|
Representation
|
|
83
91
|
--------------
|
|
84
|
-
EFAX has a single base class for its objects:
|
|
92
|
+
EFAX has a single base class for its objects: `Distribution` whose type encodes the
|
|
85
93
|
distribution family.
|
|
86
94
|
|
|
87
95
|
Each parametrization object has a shape, and so it can store any number of distributions.
|
|
@@ -89,11 +97,11 @@ Operations on these objects are vectorized.
|
|
|
89
97
|
This is unlike SciPy where each distribution is represented by a single object, and so a thousand
|
|
90
98
|
distributions need a thousand objects, and corresponding calls to functions that operate on them.
|
|
91
99
|
|
|
92
|
-
All parametrization objects are dataclasses using
|
|
93
|
-
a modification of Python's dataclasses to support JAX's
|
|
100
|
+
All parametrization objects are dataclasses using `tjax.dataclass`. These dataclasses are
|
|
101
|
+
a modification of Python's dataclasses to support JAX's "PyTree" type registration.
|
|
94
102
|
|
|
95
103
|
Each of the fields of a parametrization object stores a parameter over a specified support.
|
|
96
|
-
Some parameters are marked as
|
|
104
|
+
Some parameters are marked as "fixed", which means that they are fixed with respect to the
|
|
97
105
|
exponential family. An example of a fixed parameter is the failure number of the negative binomial
|
|
98
106
|
distribution.
|
|
99
107
|
|
|
@@ -107,17 +115,17 @@ For example:
|
|
|
107
115
|
negative_half_precision: RealArray = distribution_parameter(SymmetricMatrixSupport())
|
|
108
116
|
|
|
109
117
|
In this case, we see that there are two natural parameters for the multivariate normal distribution.
|
|
110
|
-
Objects of this type can hold any number of distributions: if such an object
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
118
|
+
Objects of this type can hold any number of distributions: if such an object `x` has shape
|
|
119
|
+
`s`, then the shape of
|
|
120
|
+
`x.mean_times_precision` is `(*s, n)` and the shape of
|
|
121
|
+
`x.negative_half_precision` is `(*s, n, n)`.
|
|
114
122
|
|
|
115
123
|
Parametrizations
|
|
116
124
|
----------------
|
|
117
125
|
Each exponential family distribution has two special parametrizations: the natural and the
|
|
118
126
|
expectation parametrization. (These are described in the overview pdf.)
|
|
119
127
|
Consequently, every distribution has at least two base classes, one inheriting from
|
|
120
|
-
|
|
128
|
+
`NaturalParametrization` and one from `ExpectationParametrization`.
|
|
121
129
|
|
|
122
130
|
The motivation for the natural parametrization is combining and scaling independent predictive
|
|
123
131
|
evidence. In the natural parametrization, these operations correspond to scaling and addition.
|
|
@@ -127,56 +135,129 @@ maximum likelihood distribution that could have produced them. In the expectati
|
|
|
127
135
|
this is an expected value.
|
|
128
136
|
|
|
129
137
|
EFAX provides conversions between the two parametrizations through the
|
|
130
|
-
|
|
138
|
+
`NaturalParametrization.to_exp` and `ExpectationParametrization.to_nat` methods.
|
|
139
|
+
|
|
140
|
+
Some distributions also provide additional convenience parametrizations. For example, the normal
|
|
141
|
+
distribution offers `NormalVP` (variance parametrization) and `NormalDP` (deviation
|
|
142
|
+
parametrization), and the multivariate normal provides `MultivariateNormalVP` and
|
|
143
|
+
`MultivariateDiagonalNormalVP`. These are not exponential-family parametrizations but are
|
|
144
|
+
useful for constructing and interpreting distributions.
|
|
131
145
|
|
|
132
146
|
Important methods
|
|
133
147
|
-----------------
|
|
134
148
|
EFAX aims to provide the main methods used in machine learning.
|
|
135
149
|
|
|
136
|
-
Every
|
|
150
|
+
Every `Distribution` has:
|
|
137
151
|
|
|
138
|
-
-
|
|
139
|
-
|
|
140
|
-
into a distribution object, operated on, and then flattened before being sent back to the model.
|
|
141
|
-
Flattening is careful with distributions with symmetric (or Hermitian) matrix-valued parameters.
|
|
142
|
-
It only stores the upper triangular elements. And,
|
|
143
|
-
- :python:`shape`, which supports broadcasting.
|
|
152
|
+
- `shape` and `ndim`, which support broadcasting, and
|
|
153
|
+
- indexing via `[]`, which slices all parameter arrays simultaneously.
|
|
144
154
|
|
|
145
|
-
Every
|
|
155
|
+
Every `NaturalParametrization` has methods:
|
|
146
156
|
|
|
147
|
-
-
|
|
148
|
-
-
|
|
157
|
+
- `to_exp` to convert itself to expectation parameters,
|
|
158
|
+
- `sufficient_statistics` to produce the sufficient statistics given an observation (used in
|
|
149
159
|
maximum likelihood estimation),
|
|
150
|
-
-
|
|
151
|
-
-
|
|
152
|
-
-
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
-
|
|
160
|
+
- `log_normalizer`, the log partition function,
|
|
161
|
+
- `carrier_measure`, the base measure,
|
|
162
|
+
- `pdf` and `log_pdf`, which are the density or mass function and its logarithm,
|
|
163
|
+
- `fisher_information_diagonal` and `fisher_information_trace`, which return the
|
|
164
|
+
diagonal and trace of the Fisher information matrix stored as distribution objects,
|
|
165
|
+
- `apply_fisher_information`, which applies the Fisher information matrix to a vector of
|
|
166
|
+
expectation parameters efficiently in a single VJP pass,
|
|
167
|
+
- `jeffreys_prior_density`, which returns the square root of the Fisher information
|
|
168
|
+
determinant,
|
|
169
|
+
- `characteristic_function`, which evaluates the characteristic function of the sufficient
|
|
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
|
|
174
|
+
- `kl_divergence`, which is the KL divergence.
|
|
175
|
+
|
|
176
|
+
Every `ExpectationParametrization` has methods:
|
|
177
|
+
|
|
178
|
+
- `to_nat` to convert itself to natural parameters, and
|
|
179
|
+
- `kl_divergence`, which is the KL divergence.
|
|
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.
|
|
158
185
|
|
|
159
186
|
Some parametrizations inherit from these interfaces:
|
|
160
187
|
|
|
161
|
-
-
|
|
162
|
-
-
|
|
163
|
-
-
|
|
164
|
-
-
|
|
188
|
+
- `HasConjugatePrior` can produce and recover the conjugate prior,
|
|
189
|
+
- `HasGeneralizedConjugatePrior` extends that with per-dimension pseudo-observation counts,
|
|
190
|
+
- `Multidimensional` distributions have an integer number of `dimensions`, and
|
|
191
|
+
- `Samplable` distributions support sampling.
|
|
165
192
|
|
|
166
193
|
Some parametrizations inherit from these public mixins:
|
|
167
194
|
|
|
168
|
-
-
|
|
169
|
-
-
|
|
170
|
-
|
|
195
|
+
- `HasEntropy` is a distribution with a `entropy` method,
|
|
196
|
+
- `HasEntropyEP` is an expectation parametrization with analytically tractable entropy and
|
|
197
|
+
`cross_entropy`, and
|
|
198
|
+
- `HasEntropyNP` is a natural parametrization with analytically tractable entropy via the
|
|
199
|
+
paired expectation parametrization.
|
|
171
200
|
|
|
172
201
|
Some parametrizations inherit from these private mixins:
|
|
173
202
|
|
|
174
|
-
-
|
|
203
|
+
- `ExpToNat` implements the conversion from expectation to natural parameters when no
|
|
175
204
|
analytical solution is possible. It uses Newton's method with a Jacobian to invert the gradient
|
|
176
205
|
log-normalizer.
|
|
177
|
-
-
|
|
206
|
+
- `TransformedNaturalParametrization` produces a natural parametrization by relating it to
|
|
178
207
|
an existing natural parametrization. And similarly for
|
|
179
|
-
|
|
208
|
+
`TransformedExpectationParametrization`.
|
|
209
|
+
|
|
210
|
+
Joint distributions
|
|
211
|
+
-------------------
|
|
212
|
+
`JointDistribution`, `JointDistributionE`, and `JointDistributionN` compose
|
|
213
|
+
multiple independent distributions into a single object. `JointDistributionE` holds
|
|
214
|
+
expectation parametrizations and implements `HasEntropyEP`; `JointDistributionN`
|
|
215
|
+
holds natural parametrizations. They support the same `to_nat` / `to_exp`
|
|
216
|
+
conversions as simple distributions.
|
|
217
|
+
|
|
218
|
+
Structure utilities
|
|
219
|
+
-------------------
|
|
220
|
+
EFAX provides three classes that capture the static metadata of a distribution tree—its types,
|
|
221
|
+
parameter names, and dimension information—without requiring a live instance. They form an
|
|
222
|
+
inheritance hierarchy:
|
|
223
|
+
|
|
224
|
+
**Assembler**
|
|
225
|
+
Stores a post-order traversal of a `Distribution` tree (types, paths, dimensions) so
|
|
226
|
+
that distributions can be reconstructed from raw parameter data without passing type information
|
|
227
|
+
alongside arrays. Key methods:
|
|
228
|
+
|
|
229
|
+
- `assemble(params)` — rebuild a `Distribution` from a `{path: array}` mapping,
|
|
230
|
+
- `coerce_from_distribution(q)` — reinterpret q's numeric values under this Assembler's types,
|
|
231
|
+
- `domain_support()` — enumerate each leaf distribution's parameter constraints,
|
|
232
|
+
- `generate_random(xp, rng, shape, safety)` — draw a random distribution with valid parameters, and
|
|
233
|
+
- `to_nat()` / `to_exp()` — return a copy whose types are all in natural or expectation form.
|
|
234
|
+
|
|
235
|
+
**Estimator** *(extends Assembler)*
|
|
236
|
+
Adds maximum likelihood estimation by recording which parameters are fixed (held constant across
|
|
237
|
+
observations) and which are free. Because the MLE for every exponential family equals the mean
|
|
238
|
+
of the sufficient statistics, estimation reduces to a single call:
|
|
239
|
+
|
|
240
|
+
- `sufficient_statistics(x)` — compute the sufficient statistics of observation x, with
|
|
241
|
+
fixed parameters supplied automatically.
|
|
242
|
+
|
|
243
|
+
Create one with `Estimator.from_type(type_p, **fixed)`,
|
|
244
|
+
`Estimator.from_expectation(p)`, or `Estimator.from_natural(p)`.
|
|
245
|
+
|
|
246
|
+
**Flattener** *(extends Estimator)*
|
|
247
|
+
Adds encoding and decoding between a `Distribution` and an array of shape
|
|
248
|
+
`(*distribution.shape, k)`, making distributions compatible with neural networks and numerical
|
|
249
|
+
optimizers. Fixed parameters are excluded from the encoded array and reinserted automatically
|
|
250
|
+
on decode.
|
|
251
|
+
|
|
252
|
+
- `Flattener.flatten(p, mapped_to_plane=True)` — encode p into a `(Flattener, array)` pair,
|
|
253
|
+
- `unflatten(array)` — decode the array back into a distribution, and
|
|
254
|
+
- `final_dimension_size()` — the size k of the last axis of the encoded array.
|
|
255
|
+
|
|
256
|
+
The `mapped_to_plane` flag controls whether constrained parameters (e.g., those on a
|
|
257
|
+
simplex or restricted to the positive reals) are bijectively mapped to all of ℝⁿ. Set it
|
|
258
|
+
`True` when passing to a neural network (to prevent invalid outputs), and `False`
|
|
259
|
+
when the raw magnitudes matter—for example when differencing expectation parameters or computing
|
|
260
|
+
Jacobians.
|
|
180
261
|
|
|
181
262
|
Distributions
|
|
182
263
|
=============
|
|
@@ -220,7 +301,7 @@ EFAX supports the following distributions:
|
|
|
220
301
|
- on a finite set:
|
|
221
302
|
|
|
222
303
|
- Bernoulli
|
|
223
|
-
-
|
|
304
|
+
- categorical
|
|
224
305
|
|
|
225
306
|
- on the nonnegative integers:
|
|
226
307
|
|
|
@@ -246,9 +327,14 @@ EFAX supports the following distributions:
|
|
|
246
327
|
- Dirichlet
|
|
247
328
|
- generalized Dirichlet
|
|
248
329
|
|
|
249
|
-
- on
|
|
330
|
+
- on spheres:
|
|
250
331
|
|
|
251
332
|
- von Mises-Fisher
|
|
333
|
+
- complex von Mises
|
|
334
|
+
|
|
335
|
+
- on positive-definite matrices:
|
|
336
|
+
|
|
337
|
+
- Wishart
|
|
252
338
|
|
|
253
339
|
Usage
|
|
254
340
|
=====
|
|
@@ -376,17 +462,23 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
376
462
|
return x - 1e-4 * x_bar
|
|
377
463
|
|
|
378
464
|
|
|
379
|
-
|
|
380
|
-
|
|
381
|
-
|
|
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]
|
|
382
468
|
|
|
383
469
|
|
|
384
|
-
def cond_fun(
|
|
385
|
-
q_bar =
|
|
470
|
+
def cond_fun(state: State) -> JaxBooleanArray:
|
|
471
|
+
_, q_bar = state
|
|
386
472
|
total = jnp.sum(parameter_dot_product(q_bar, q_bar))
|
|
387
473
|
return total > 1e-6 # noqa: PLR2004
|
|
388
474
|
|
|
389
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
|
+
|
|
390
482
|
# The target_distribution is represented as the expectation parameters of a
|
|
391
483
|
# Bernoulli distribution corresponding to probabilities 0.3, 0.4, and 0.7.
|
|
392
484
|
target_distribution = BernoulliEP(jnp.asarray([0.3, 0.4, 0.7]))
|
|
@@ -395,10 +487,13 @@ Using the cross entropy to iteratively optimize a prediction is simple:
|
|
|
395
487
|
# of a Bernoulli distribution corresponding to log-odds 0, which is probability
|
|
396
488
|
# 0.5.
|
|
397
489
|
initial_predictive_distribution = BernoulliNP(jnp.zeros(3))
|
|
490
|
+
initial_gradient = gradient_cross_entropy(target_distribution,
|
|
491
|
+
initial_predictive_distribution)
|
|
398
492
|
|
|
399
493
|
# Optimize the predictive distribution iteratively.
|
|
400
|
-
predictive_distribution = lax.while_loop(
|
|
401
|
-
|
|
494
|
+
predictive_distribution, _ = lax.while_loop(
|
|
495
|
+
cond_fun, body_fun, (initial_predictive_distribution, initial_gradient)
|
|
496
|
+
)
|
|
402
497
|
|
|
403
498
|
# Compare the optimized predictive distribution with the target value in the
|
|
404
499
|
# same natural parametrization.
|
|
@@ -437,14 +532,14 @@ instead.
|
|
|
437
532
|
likelihood estimation.
|
|
438
533
|
|
|
439
534
|
Suppose you have some samples from a distribution family with unknown
|
|
440
|
-
parameters, and you want to estimate the maximum likelihood
|
|
535
|
+
parameters, and you want to estimate the maximum likelihood parameters of the
|
|
441
536
|
distribution.
|
|
442
537
|
"""
|
|
443
538
|
import jax.numpy as jnp
|
|
444
539
|
import jax.random as jr
|
|
445
540
|
from tjax import print_generic
|
|
446
541
|
|
|
447
|
-
from efax import DirichletEP, DirichletNP,
|
|
542
|
+
from efax import DirichletEP, DirichletNP, Estimator, parameter_mean
|
|
448
543
|
|
|
449
544
|
# Consider a Dirichlet distribution with a given alpha.
|
|
450
545
|
alpha = jnp.asarray([2.0, 3.0, 4.0])
|
|
@@ -457,12 +552,14 @@ instead.
|
|
|
457
552
|
|
|
458
553
|
# Now, let's find the maximum likelihood Dirichlet distribution that fits it.
|
|
459
554
|
# First, convert the samples to their sufficient statistics.
|
|
460
|
-
estimator =
|
|
555
|
+
estimator = Estimator.from_type(DirichletEP)
|
|
461
556
|
ss = estimator.sufficient_statistics(samples)
|
|
462
|
-
# ss has type DirichletEP. This is similar to the
|
|
557
|
+
# ss has type DirichletEP. This is similar to the conjugate prior of the
|
|
463
558
|
# Dirichlet distribution.
|
|
464
559
|
|
|
465
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.
|
|
466
563
|
ss_mean = parameter_mean(ss, axis=0) # ss_mean also has type DirichletEP.
|
|
467
564
|
|
|
468
565
|
# Convert this back to the natural parametrization.
|