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.
Files changed (130) hide show
  1. {efax-2.1.2 → efax-2.2.3}/PKG-INFO +156 -59
  2. {efax-2.1.2 → efax-2.2.3}/README.rst +148 -51
  3. {efax-2.1.2 → efax-2.2.3}/efax/__init__.py +7 -7
  4. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/bernoulli.py +2 -1
  5. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/beta.py +1 -0
  6. efax-2.1.2/efax/_src/distributions/multinomial.py → efax-2.2.3/efax/_src/distributions/categorical.py +20 -18
  7. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/chi.py +2 -1
  8. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/unit_variance.py +1 -1
  9. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/unit_variance.py +1 -1
  10. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/dirichlet.py +1 -0
  11. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/exponential.py +2 -1
  12. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/gamma.py +8 -4
  13. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/geometric.py +1 -0
  14. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/inverse_gamma.py +2 -1
  15. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/log_normal.py +2 -2
  16. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/unit_variance.py +4 -4
  17. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/fixed_variance.py +2 -1
  18. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/unit_variance.py +2 -1
  19. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/negative_binomial.py +1 -0
  20. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/unit_variance.py +2 -1
  21. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/poisson.py +2 -1
  22. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/rayleigh.py +2 -1
  23. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/softplus.py +2 -2
  24. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/unit_variance.py +4 -4
  25. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/von_mises.py +3 -3
  26. {efax-2.1.2 → efax-2.2.3}/efax/_src/expectation_parametrization.py +7 -6
  27. efax-2.2.3/efax/_src/interfaces/conjugate_prior.py +72 -0
  28. efax-2.2.3/efax/_src/interfaces/multidimensional.py +18 -0
  29. {efax-2.1.2 → efax-2.2.3}/efax/_src/interfaces/samplable.py +9 -0
  30. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +12 -3
  31. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -1
  32. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/has_entropy.py +23 -8
  33. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/transformed_parametrization.py +36 -9
  34. {efax-2.1.2 → efax-2.2.3}/efax/_src/natural_parametrization.py +43 -18
  35. {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/support.py +7 -1
  36. {efax-2.1.2 → efax-2.2.3}/efax/_src/parametrization.py +24 -5
  37. efax-2.1.2/efax/_src/structure/structure.py → efax-2.2.3/efax/_src/structure/assembler.py +54 -26
  38. {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/estimator.py +57 -25
  39. {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/flattener.py +46 -34
  40. {efax-2.1.2 → efax-2.2.3}/efax/_src/tools.py +10 -13
  41. {efax-2.1.2 → efax-2.2.3}/efax/_src/transform/joint.py +13 -3
  42. {efax-2.1.2 → efax-2.2.3}/pyproject.toml +8 -9
  43. efax-2.1.2/.editorconfig +0 -19
  44. efax-2.1.2/.gitignore +0 -2
  45. efax-2.1.2/LICENSE +0 -201
  46. efax-2.1.2/efax/_src/interfaces/conjugate_prior.py +0 -53
  47. efax-2.1.2/efax/_src/interfaces/multidimensional.py +0 -11
  48. efax-2.1.2/examples/.editorconfig +0 -2
  49. efax-2.1.2/examples/__init__.py +0 -1
  50. efax-2.1.2/examples/bayesian_evidence_combination.py +0 -48
  51. efax-2.1.2/examples/cross_entropy.py +0 -33
  52. efax-2.1.2/examples/maximum_likelihood_estimation.py +0 -47
  53. efax-2.1.2/examples/optimization.py +0 -79
  54. efax-2.1.2/exponential_families.pdf +3 -29318
  55. efax-2.1.2/lefthook.yml +0 -18
  56. efax-2.1.2/tests/__init__.py +0 -1
  57. efax-2.1.2/tests/conftest.py +0 -152
  58. efax-2.1.2/tests/create_info.py +0 -857
  59. efax-2.1.2/tests/distribution_info.py +0 -116
  60. efax-2.1.2/tests/match_scipy/__init__.py +0 -1
  61. efax-2.1.2/tests/match_scipy/test_entropy.py +0 -33
  62. efax-2.1.2/tests/match_scipy/test_maximum_likelihood_estimation.py +0 -78
  63. efax-2.1.2/tests/match_scipy/test_pdf.py +0 -71
  64. efax-2.1.2/tests/scipy_replacement/__init__.py +0 -1
  65. efax-2.1.2/tests/scipy_replacement/base.py +0 -24
  66. efax-2.1.2/tests/scipy_replacement/complex_multivariate_normal.py +0 -160
  67. efax-2.1.2/tests/scipy_replacement/complex_normal.py +0 -157
  68. efax-2.1.2/tests/scipy_replacement/dirichlet.py +0 -92
  69. efax-2.1.2/tests/scipy_replacement/joint.py +0 -33
  70. efax-2.1.2/tests/scipy_replacement/multinomial.py +0 -26
  71. efax-2.1.2/tests/scipy_replacement/multivariate_normal.py +0 -54
  72. efax-2.1.2/tests/scipy_replacement/shaped_distribution.py +0 -92
  73. efax-2.1.2/tests/scipy_replacement/von_mises.py +0 -44
  74. efax-2.1.2/tests/scipy_replacement/wishart.py +0 -32
  75. efax-2.1.2/tests/softplus.py +0 -16
  76. efax-2.1.2/tests/test_characteristic_function.py +0 -138
  77. efax-2.1.2/tests/test_complex_normal.py +0 -105
  78. efax-2.1.2/tests/test_complex_von_mises.py +0 -27
  79. efax-2.1.2/tests/test_conjugate_prior.py +0 -95
  80. efax-2.1.2/tests/test_conversion.py +0 -42
  81. efax-2.1.2/tests/test_degenerate.py +0 -38
  82. efax-2.1.2/tests/test_entropy_gradient.py +0 -90
  83. efax-2.1.2/tests/test_ep_from_cf.py +0 -127
  84. efax-2.1.2/tests/test_fisher_information.py +0 -93
  85. efax-2.1.2/tests/test_flatten.py +0 -42
  86. efax-2.1.2/tests/test_gradient_log_normalizer.py +0 -111
  87. efax-2.1.2/tests/test_hessian.py +0 -97
  88. efax-2.1.2/tests/test_jax_quirks.py +0 -14
  89. efax-2.1.2/tests/test_kl.py +0 -85
  90. efax-2.1.2/tests/test_reparametrization_trick.py +0 -22
  91. efax-2.1.2/tests/test_sampling.py +0 -193
  92. efax-2.1.2/tests/test_scipy_distributions.py +0 -32
  93. efax-2.1.2/tests/test_shapes.py +0 -38
  94. efax-2.1.2/uv.lock +0 -2199
  95. {efax-2.1.2 → efax-2.2.3}/efax/_src/__init__.py +0 -0
  96. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/__init__.py +0 -0
  97. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/chi_square.py +0 -0
  98. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/__init__.py +0 -0
  99. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
  100. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/__init__.py +0 -0
  101. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
  102. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/complex_von_mises.py +1 -1
  103. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/dirichlet_common.py +0 -0
  104. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/gen_dirichlet.py +0 -0
  105. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/inverse_gaussian.py +0 -0
  106. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/log_normal/__init__.py +0 -0
  107. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/logarithmic.py +0 -0
  108. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
  109. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
  110. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
  111. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
  112. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/negative_binomial_common.py +0 -0
  113. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/__init__.py +0 -0
  114. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/normal/normal.py +0 -0
  115. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
  116. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/weibull.py +0 -0
  117. {efax-2.1.2 → efax-2.2.3}/efax/_src/distributions/wishart.py +1 -1
  118. {efax-2.1.2 → efax-2.2.3}/efax/_src/interfaces/__init__.py +0 -0
  119. {efax-2.1.2 → efax-2.2.3}/efax/_src/iteration.py +0 -0
  120. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/__init__.py +0 -0
  121. {efax-2.1.2 → efax-2.2.3}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
  122. {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/__init__.py +0 -0
  123. {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/parameter.py +0 -0
  124. {efax-2.1.2 → efax-2.2.3}/efax/_src/parameter/ring.py +0 -0
  125. {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/__init__.py +0 -0
  126. {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/parameter_names.py +0 -0
  127. {efax-2.1.2 → efax-2.2.3}/efax/_src/structure/parameter_supports.py +0 -0
  128. {efax-2.1.2 → efax-2.2.3}/efax/_src/transform/__init__.py +0 -0
  129. {efax-2.1.2 → efax-2.2.3}/efax/_src/types.py +0 -0
  130. {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.1.2
3
+ Version: 2.2.3
4
4
  Summary: Exponential families for JAX
5
- Project-URL: source, https://github.com/NeilGirdhar/efax
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.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: :python:`Distribution` whose type encodes the
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 :python:`tjax.dataclass`. These dataclasses are
93
- a modification of Python's dataclasses to support JAX's “PyTree” type registration.
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 “fixed”, which means that they are fixed with respect to the
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 :python:`x` has shape
111
- :python:`s`, then the shape of
112
- :python:`x.mean_times_precision` is :python:`(*s, n)` and the shape of
113
- :python:`x.negative_half_precision` is :python:`(*s, n, n)`.
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
- :python:`NaturalParametrization` and one from :python:`ExpectationParametrization`.
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
- :python:`NaturalParametrization.to_exp` and :python:`ExpectationParametrization.to_nat` methods.
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 :python:`Distribution` has methods:
150
+ Every `Distribution` has:
137
151
 
138
- - :python:`flattened` and :python:`unflattened` to flatten and unflatten the parameters into a
139
- single array. Typically, array-valued signals in a machine learning model would be unflattened
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 :python:`NaturalParametrization` has methods:
155
+ Every `NaturalParametrization` has methods:
146
156
 
147
- - :python:`to_exp` to convert itself to expectation parameters.
148
- - :python:`sufficient_statistics` to produce the sufficient statistics given an observation (used in
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
- - :python:`pdf` and :python:`log_pdf`, which is the density or mass function and its logarithm,
151
- - :python:`fisher_information`, which is the Fisher information matrix, and
152
- - :python:`kl_divergence`, which is the KL divergence.
153
-
154
- Every :python:`ExpectationParametrization` has methods:
155
-
156
- - :python:`to_nat` to convert itself to natural parameters, and
157
- - :python:`kl_divergence`, which is the KL divergence.
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
- - :python:`HasConjugatePrior` can produce the conjugate prior,
162
- - :python:`HasGeneralizedConjugatePrior` can produce a generalization of the conjugate prior,
163
- - :python:`Multidimensional` distributions have a integer number of `dimensions`, and
164
- - :python:`Samplable` distributions support sampling.
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
- - :python:`HasEntropyEP` is an expectation parametrization with an entropy and cross entropy, and
169
- - :python:`HasEntropyNP` is a natural parametrization with an entropy, (The cross entropy is not
170
- efficient.)
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
- - :python:`ExpToNat` implements the conversion from expectation to natural parameters when no
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
- - :python:`TransformedNaturalParametrization` produces a natural parametrization by relating it to
206
+ - `TransformedNaturalParametrization` produces a natural parametrization by relating it to
178
207
  an existing natural parametrization. And similarly for
179
- :python:`TransformedExpectationParametrization`.
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
- - multinomial
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 the n-sphere:
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
- def body_fun(q: BernoulliNP) -> BernoulliNP:
380
- q_bar = gradient_cross_entropy(target_distribution, q)
381
- return parameter_map(apply, q, q_bar)
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(q: BernoulliNP) -> JaxBooleanArray:
385
- q_bar = gradient_cross_entropy(target_distribution, q)
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(cond_fun, body_fun,
401
- initial_predictive_distribution)
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 parmaters of the
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, MaximumLikelihoodEstimator, parameter_mean
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 = MaximumLikelihoodEstimator.create_simple_estimator(DirichletEP)
555
+ estimator = Estimator.from_type(DirichletEP)
461
556
  ss = estimator.sufficient_statistics(samples)
462
- # ss has type DirichletEP. This is similar to the conjguate prior of 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.