efax 2.1.2__tar.gz → 2.2.0__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.0}/PKG-INFO +122 -50
  2. {efax-2.1.2 → efax-2.2.0}/README.rst +115 -43
  3. {efax-2.1.2 → efax-2.2.0}/efax/__init__.py +7 -7
  4. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/bernoulli.py +2 -1
  5. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/beta.py +1 -0
  6. efax-2.1.2/efax/_src/distributions/multinomial.py → efax-2.2.0/efax/_src/distributions/categorical.py +19 -18
  7. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/chi.py +2 -1
  8. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/unit_variance.py +1 -1
  9. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/unit_variance.py +1 -1
  10. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/dirichlet.py +1 -0
  11. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/exponential.py +2 -1
  12. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/gamma.py +2 -1
  13. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/geometric.py +1 -0
  14. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/inverse_gamma.py +2 -1
  15. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/log_normal.py +2 -2
  16. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/unit_variance.py +4 -4
  17. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/fixed_variance.py +2 -1
  18. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/unit_variance.py +2 -1
  19. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/negative_binomial.py +1 -0
  20. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/unit_variance.py +2 -1
  21. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/poisson.py +2 -1
  22. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/rayleigh.py +2 -1
  23. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/softplus.py +2 -2
  24. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/unit_variance.py +4 -4
  25. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/von_mises.py +3 -3
  26. {efax-2.1.2 → efax-2.2.0}/efax/_src/expectation_parametrization.py +7 -6
  27. efax-2.2.0/efax/_src/interfaces/conjugate_prior.py +72 -0
  28. efax-2.2.0/efax/_src/interfaces/multidimensional.py +18 -0
  29. {efax-2.1.2 → efax-2.2.0}/efax/_src/interfaces/samplable.py +9 -0
  30. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +12 -3
  31. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/has_entropy.py +23 -8
  32. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/transformed_parametrization.py +36 -9
  33. {efax-2.1.2 → efax-2.2.0}/efax/_src/natural_parametrization.py +18 -12
  34. {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/support.py +1 -0
  35. {efax-2.1.2 → efax-2.2.0}/efax/_src/parametrization.py +24 -5
  36. efax-2.1.2/efax/_src/structure/structure.py → efax-2.2.0/efax/_src/structure/assembler.py +47 -23
  37. {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/estimator.py +51 -23
  38. {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/flattener.py +46 -34
  39. {efax-2.1.2 → efax-2.2.0}/efax/_src/tools.py +3 -3
  40. {efax-2.1.2 → efax-2.2.0}/pyproject.toml +6 -6
  41. efax-2.1.2/.editorconfig +0 -19
  42. efax-2.1.2/.gitignore +0 -2
  43. efax-2.1.2/LICENSE +0 -201
  44. efax-2.1.2/efax/_src/interfaces/conjugate_prior.py +0 -53
  45. efax-2.1.2/efax/_src/interfaces/multidimensional.py +0 -11
  46. efax-2.1.2/examples/.editorconfig +0 -2
  47. efax-2.1.2/examples/__init__.py +0 -1
  48. efax-2.1.2/examples/bayesian_evidence_combination.py +0 -48
  49. efax-2.1.2/examples/cross_entropy.py +0 -33
  50. efax-2.1.2/examples/maximum_likelihood_estimation.py +0 -47
  51. efax-2.1.2/examples/optimization.py +0 -79
  52. efax-2.1.2/exponential_families.pdf +3 -29318
  53. efax-2.1.2/lefthook.yml +0 -18
  54. efax-2.1.2/tests/__init__.py +0 -1
  55. efax-2.1.2/tests/conftest.py +0 -152
  56. efax-2.1.2/tests/create_info.py +0 -857
  57. efax-2.1.2/tests/distribution_info.py +0 -116
  58. efax-2.1.2/tests/match_scipy/__init__.py +0 -1
  59. efax-2.1.2/tests/match_scipy/test_entropy.py +0 -33
  60. efax-2.1.2/tests/match_scipy/test_maximum_likelihood_estimation.py +0 -78
  61. efax-2.1.2/tests/match_scipy/test_pdf.py +0 -71
  62. efax-2.1.2/tests/scipy_replacement/__init__.py +0 -1
  63. efax-2.1.2/tests/scipy_replacement/base.py +0 -24
  64. efax-2.1.2/tests/scipy_replacement/complex_multivariate_normal.py +0 -160
  65. efax-2.1.2/tests/scipy_replacement/complex_normal.py +0 -157
  66. efax-2.1.2/tests/scipy_replacement/dirichlet.py +0 -92
  67. efax-2.1.2/tests/scipy_replacement/joint.py +0 -33
  68. efax-2.1.2/tests/scipy_replacement/multinomial.py +0 -26
  69. efax-2.1.2/tests/scipy_replacement/multivariate_normal.py +0 -54
  70. efax-2.1.2/tests/scipy_replacement/shaped_distribution.py +0 -92
  71. efax-2.1.2/tests/scipy_replacement/von_mises.py +0 -44
  72. efax-2.1.2/tests/scipy_replacement/wishart.py +0 -32
  73. efax-2.1.2/tests/softplus.py +0 -16
  74. efax-2.1.2/tests/test_characteristic_function.py +0 -138
  75. efax-2.1.2/tests/test_complex_normal.py +0 -105
  76. efax-2.1.2/tests/test_complex_von_mises.py +0 -27
  77. efax-2.1.2/tests/test_conjugate_prior.py +0 -95
  78. efax-2.1.2/tests/test_conversion.py +0 -42
  79. efax-2.1.2/tests/test_degenerate.py +0 -38
  80. efax-2.1.2/tests/test_entropy_gradient.py +0 -90
  81. efax-2.1.2/tests/test_ep_from_cf.py +0 -127
  82. efax-2.1.2/tests/test_fisher_information.py +0 -93
  83. efax-2.1.2/tests/test_flatten.py +0 -42
  84. efax-2.1.2/tests/test_gradient_log_normalizer.py +0 -111
  85. efax-2.1.2/tests/test_hessian.py +0 -97
  86. efax-2.1.2/tests/test_jax_quirks.py +0 -14
  87. efax-2.1.2/tests/test_kl.py +0 -85
  88. efax-2.1.2/tests/test_reparametrization_trick.py +0 -22
  89. efax-2.1.2/tests/test_sampling.py +0 -193
  90. efax-2.1.2/tests/test_scipy_distributions.py +0 -32
  91. efax-2.1.2/tests/test_shapes.py +0 -38
  92. efax-2.1.2/uv.lock +0 -2199
  93. {efax-2.1.2 → efax-2.2.0}/efax/_src/__init__.py +0 -0
  94. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/__init__.py +0 -0
  95. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/chi_square.py +0 -0
  96. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/__init__.py +0 -0
  97. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
  98. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/__init__.py +0 -0
  99. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
  100. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_von_mises.py +1 -1
  101. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/dirichlet_common.py +0 -0
  102. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/gen_dirichlet.py +0 -0
  103. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/inverse_gaussian.py +0 -0
  104. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/__init__.py +0 -0
  105. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/logarithmic.py +0 -0
  106. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
  107. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
  108. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
  109. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
  110. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/negative_binomial_common.py +0 -0
  111. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/__init__.py +0 -0
  112. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/normal.py +0 -0
  113. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
  114. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/weibull.py +0 -0
  115. {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/wishart.py +1 -1
  116. {efax-2.1.2 → efax-2.2.0}/efax/_src/interfaces/__init__.py +0 -0
  117. {efax-2.1.2 → efax-2.2.0}/efax/_src/iteration.py +0 -0
  118. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/__init__.py +0 -0
  119. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
  120. {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/optimistix.py +0 -0
  121. {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/__init__.py +0 -0
  122. {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/parameter.py +0 -0
  123. {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/ring.py +0 -0
  124. {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/__init__.py +0 -0
  125. {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/parameter_names.py +0 -0
  126. {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/parameter_supports.py +0 -0
  127. {efax-2.1.2 → efax-2.2.0}/efax/_src/transform/__init__.py +0 -0
  128. {efax-2.1.2 → efax-2.2.0}/efax/_src/transform/joint.py +0 -0
  129. {efax-2.1.2 → efax-2.2.0}/efax/_src/types.py +0 -0
  130. {efax-2.1.2 → efax-2.2.0}/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.0
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
@@ -29,6 +25,10 @@ Requires-Dist: optimistix>=0.0.9
29
25
  Requires-Dist: optype[numpy]>=0.8.0
30
26
  Requires-Dist: tjax>=1.4.4
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
@@ -81,7 +83,7 @@ Framework
81
83
  =========
82
84
  Representation
83
85
  --------------
84
- EFAX has a single base class for its objects: :python:`Distribution` whose type encodes the
86
+ EFAX has a single base class for its objects: `Distribution` whose type encodes the
85
87
  distribution family.
86
88
 
87
89
  Each parametrization object has a shape, and so it can store any number of distributions.
@@ -89,11 +91,11 @@ Operations on these objects are vectorized.
89
91
  This is unlike SciPy where each distribution is represented by a single object, and so a thousand
90
92
  distributions need a thousand objects, and corresponding calls to functions that operate on them.
91
93
 
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.
94
+ All parametrization objects are dataclasses using `tjax.dataclass`. These dataclasses are
95
+ a modification of Python's dataclasses to support JAX's "PyTree" type registration.
94
96
 
95
97
  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
98
+ Some parameters are marked as "fixed", which means that they are fixed with respect to the
97
99
  exponential family. An example of a fixed parameter is the failure number of the negative binomial
98
100
  distribution.
99
101
 
@@ -107,17 +109,17 @@ For example:
107
109
  negative_half_precision: RealArray = distribution_parameter(SymmetricMatrixSupport())
108
110
 
109
111
  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)`.
112
+ Objects of this type can hold any number of distributions: if such an object `x` has shape
113
+ `s`, then the shape of
114
+ `x.mean_times_precision` is `(*s, n)` and the shape of
115
+ `x.negative_half_precision` is `(*s, n, n)`.
114
116
 
115
117
  Parametrizations
116
118
  ----------------
117
119
  Each exponential family distribution has two special parametrizations: the natural and the
118
120
  expectation parametrization. (These are described in the overview pdf.)
119
121
  Consequently, every distribution has at least two base classes, one inheriting from
120
- :python:`NaturalParametrization` and one from :python:`ExpectationParametrization`.
122
+ `NaturalParametrization` and one from `ExpectationParametrization`.
121
123
 
122
124
  The motivation for the natural parametrization is combining and scaling independent predictive
123
125
  evidence. In the natural parametrization, these operations correspond to scaling and addition.
@@ -127,56 +129,121 @@ maximum likelihood distribution that could have produced them. In the expectati
127
129
  this is an expected value.
128
130
 
129
131
  EFAX provides conversions between the two parametrizations through the
130
- :python:`NaturalParametrization.to_exp` and :python:`ExpectationParametrization.to_nat` methods.
132
+ `NaturalParametrization.to_exp` and `ExpectationParametrization.to_nat` methods.
133
+
134
+ Some distributions also provide additional convenience parametrizations. For example, the normal
135
+ distribution offers `NormalVP` (variance parametrization) and `NormalDP` (deviation
136
+ parametrization), and the multivariate normal provides `MultivariateNormalVP` and
137
+ `MultivariateDiagonalNormalVP`. These are not exponential-family parametrizations but are
138
+ useful for constructing and interpreting distributions.
131
139
 
132
140
  Important methods
133
141
  -----------------
134
142
  EFAX aims to provide the main methods used in machine learning.
135
143
 
136
- Every :python:`Distribution` has methods:
144
+ Every `Distribution` has:
137
145
 
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.
146
+ - `shape` and `ndim`, which support broadcasting, and
147
+ - indexing via `[]`, which slices all parameter arrays simultaneously.
144
148
 
145
- Every :python:`NaturalParametrization` has methods:
149
+ Every `NaturalParametrization` has methods:
146
150
 
147
- - :python:`to_exp` to convert itself to expectation parameters.
148
- - :python:`sufficient_statistics` to produce the sufficient statistics given an observation (used in
151
+ - `to_exp` to convert itself to expectation parameters,
152
+ - `sufficient_statistics` to produce the sufficient statistics given an observation (used in
149
153
  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.
154
+ - `log_normalizer`, the log partition function,
155
+ - `carrier_measure`, the base measure,
156
+ - `pdf` and `log_pdf`, which are the density or mass function and its logarithm,
157
+ - `fisher_information_diagonal` and `fisher_information_trace`, which return the
158
+ diagonal and trace of the Fisher information matrix stored as distribution objects,
159
+ - `apply_fisher_information`, which applies the Fisher information matrix to a vector of
160
+ expectation parameters efficiently in a single VJP pass,
161
+ - `jeffreys_prior_density`, which returns the square root of the Fisher information
162
+ determinant,
163
+ - `characteristic_function`, which evaluates the characteristic function of the sufficient
164
+ statistics via analytic continuation of the log-normalizer, and
165
+ - `kl_divergence`, which is the KL divergence.
166
+
167
+ Every `ExpectationParametrization` has methods:
168
+
169
+ - `to_nat` to convert itself to natural parameters, and
170
+ - `kl_divergence`, which is the KL divergence.
158
171
 
159
172
  Some parametrizations inherit from these interfaces:
160
173
 
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.
174
+ - `HasConjugatePrior` can produce and recover the conjugate prior,
175
+ - `HasGeneralizedConjugatePrior` extends that with per-dimension pseudo-observation counts,
176
+ - `Multidimensional` distributions have an integer number of `dimensions`, and
177
+ - `Samplable` distributions support sampling.
165
178
 
166
179
  Some parametrizations inherit from these public mixins:
167
180
 
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.)
181
+ - `HasEntropy` is a distribution with a `entropy` method,
182
+ - `HasEntropyEP` is an expectation parametrization with analytically tractable entropy and
183
+ `cross_entropy`, and
184
+ - `HasEntropyNP` is a natural parametrization with analytically tractable entropy via the
185
+ paired expectation parametrization.
171
186
 
172
187
  Some parametrizations inherit from these private mixins:
173
188
 
174
- - :python:`ExpToNat` implements the conversion from expectation to natural parameters when no
189
+ - `ExpToNat` implements the conversion from expectation to natural parameters when no
175
190
  analytical solution is possible. It uses Newton's method with a Jacobian to invert the gradient
176
191
  log-normalizer.
177
- - :python:`TransformedNaturalParametrization` produces a natural parametrization by relating it to
192
+ - `TransformedNaturalParametrization` produces a natural parametrization by relating it to
178
193
  an existing natural parametrization. And similarly for
179
- :python:`TransformedExpectationParametrization`.
194
+ `TransformedExpectationParametrization`.
195
+
196
+ Joint distributions
197
+ -------------------
198
+ `JointDistribution`, `JointDistributionE`, and `JointDistributionN` compose
199
+ multiple independent distributions into a single object. `JointDistributionE` holds
200
+ expectation parametrizations and implements `HasEntropyEP`; `JointDistributionN`
201
+ holds natural parametrizations. They support the same `to_nat` / `to_exp`
202
+ conversions as simple distributions.
203
+
204
+ Structure utilities
205
+ -------------------
206
+ EFAX provides three classes that capture the static metadata of a distribution tree—its types,
207
+ parameter names, and dimension information—without requiring a live instance. They form an
208
+ inheritance hierarchy:
209
+
210
+ **Assembler**
211
+ Stores a post-order traversal of a `Distribution` tree (types, paths, dimensions) so
212
+ that distributions can be reconstructed from raw parameter data without passing type information
213
+ alongside arrays. Key methods:
214
+
215
+ - `assemble(params)` — rebuild a `Distribution` from a `{path: array}` mapping,
216
+ - `coerce_from_distribution(q)` — reinterpret q's numeric values under this Assembler's types,
217
+ - `domain_support()` — enumerate each leaf distribution's parameter constraints,
218
+ - `generate_random(xp, rng, shape, safety)` — draw a random distribution with valid parameters, and
219
+ - `to_nat()` / `to_exp()` — return a copy whose types are all in natural or expectation form.
220
+
221
+ **Estimator** *(extends Assembler)*
222
+ Adds maximum likelihood estimation by recording which parameters are fixed (held constant across
223
+ observations) and which are free. Because the MLE for every exponential family equals the mean
224
+ of the sufficient statistics, estimation reduces to a single call:
225
+
226
+ - `sufficient_statistics(x)` — compute the sufficient statistics of observation x, with
227
+ fixed parameters supplied automatically.
228
+
229
+ Create one with `Estimator.from_type(type_p, **fixed)`,
230
+ `Estimator.from_expectation(p)`, or `Estimator.from_natural(p)`.
231
+
232
+ **Flattener** *(extends Estimator)*
233
+ Adds encoding and decoding between a `Distribution` and an array of shape
234
+ `(*distribution.shape, k)`, making distributions compatible with neural networks and numerical
235
+ optimizers. Fixed parameters are excluded from the encoded array and reinserted automatically
236
+ on decode.
237
+
238
+ - `Flattener.flatten(p, mapped_to_plane=True)` — encode p into a `(Flattener, array)` pair,
239
+ - `unflatten(array)` — decode the array back into a distribution, and
240
+ - `final_dimension_size()` — the size k of the last axis of the encoded array.
241
+
242
+ The `mapped_to_plane` flag controls whether constrained parameters (e.g., those on a
243
+ simplex or restricted to the positive reals) are bijectively mapped to all of ℝⁿ. Set it
244
+ `True` when passing to a neural network (to prevent invalid outputs), and `False`
245
+ when the raw magnitudes matter—for example when differencing expectation parameters or computing
246
+ Jacobians.
180
247
 
181
248
  Distributions
182
249
  =============
@@ -220,7 +287,7 @@ EFAX supports the following distributions:
220
287
  - on a finite set:
221
288
 
222
289
  - Bernoulli
223
- - multinomial
290
+ - categorical
224
291
 
225
292
  - on the nonnegative integers:
226
293
 
@@ -249,6 +316,11 @@ EFAX supports the following distributions:
249
316
  - on the n-sphere:
250
317
 
251
318
  - von Mises-Fisher
319
+ - complex von Mises
320
+
321
+ - on positive-definite matrices:
322
+
323
+ - Wishart
252
324
 
253
325
  Usage
254
326
  =====
@@ -437,14 +509,14 @@ instead.
437
509
  likelihood estimation.
438
510
 
439
511
  Suppose you have some samples from a distribution family with unknown
440
- parameters, and you want to estimate the maximum likelihood parmaters of the
512
+ parameters, and you want to estimate the maximum likelihood parameters of the
441
513
  distribution.
442
514
  """
443
515
  import jax.numpy as jnp
444
516
  import jax.random as jr
445
517
  from tjax import print_generic
446
518
 
447
- from efax import DirichletEP, DirichletNP, MaximumLikelihoodEstimator, parameter_mean
519
+ from efax import DirichletEP, DirichletNP, Estimator, parameter_mean
448
520
 
449
521
  # Consider a Dirichlet distribution with a given alpha.
450
522
  alpha = jnp.asarray([2.0, 3.0, 4.0])
@@ -457,9 +529,9 @@ instead.
457
529
 
458
530
  # Now, let's find the maximum likelihood Dirichlet distribution that fits it.
459
531
  # First, convert the samples to their sufficient statistics.
460
- estimator = MaximumLikelihoodEstimator.create_simple_estimator(DirichletEP)
532
+ estimator = Estimator.from_type(DirichletEP)
461
533
  ss = estimator.sufficient_statistics(samples)
462
- # ss has type DirichletEP. This is similar to the conjguate prior of the
534
+ # ss has type DirichletEP. This is similar to the conjugate prior of the
463
535
  # Dirichlet distribution.
464
536
 
465
537
  # Take the mean over the first axis.
@@ -4,6 +4,8 @@
4
4
  .. role:: python(code)
5
5
  :language: python
6
6
 
7
+ .. default-role:: python
8
+
7
9
  .. image:: https://img.shields.io/pypi/v/efax
8
10
  :target: https://pypi.org/project/efax/
9
11
  :alt: PyPI - Version
@@ -48,7 +50,7 @@ Framework
48
50
  =========
49
51
  Representation
50
52
  --------------
51
- EFAX has a single base class for its objects: :python:`Distribution` whose type encodes the
53
+ EFAX has a single base class for its objects: `Distribution` whose type encodes the
52
54
  distribution family.
53
55
 
54
56
  Each parametrization object has a shape, and so it can store any number of distributions.
@@ -56,11 +58,11 @@ Operations on these objects are vectorized.
56
58
  This is unlike SciPy where each distribution is represented by a single object, and so a thousand
57
59
  distributions need a thousand objects, and corresponding calls to functions that operate on them.
58
60
 
59
- All parametrization objects are dataclasses using :python:`tjax.dataclass`. These dataclasses are
60
- a modification of Python's dataclasses to support JAX's “PyTree” type registration.
61
+ All parametrization objects are dataclasses using `tjax.dataclass`. These dataclasses are
62
+ a modification of Python's dataclasses to support JAX's "PyTree" type registration.
61
63
 
62
64
  Each of the fields of a parametrization object stores a parameter over a specified support.
63
- Some parameters are marked as “fixed”, which means that they are fixed with respect to the
65
+ Some parameters are marked as "fixed", which means that they are fixed with respect to the
64
66
  exponential family. An example of a fixed parameter is the failure number of the negative binomial
65
67
  distribution.
66
68
 
@@ -74,17 +76,17 @@ For example:
74
76
  negative_half_precision: RealArray = distribution_parameter(SymmetricMatrixSupport())
75
77
 
76
78
  In this case, we see that there are two natural parameters for the multivariate normal distribution.
77
- Objects of this type can hold any number of distributions: if such an object :python:`x` has shape
78
- :python:`s`, then the shape of
79
- :python:`x.mean_times_precision` is :python:`(*s, n)` and the shape of
80
- :python:`x.negative_half_precision` is :python:`(*s, n, n)`.
79
+ Objects of this type can hold any number of distributions: if such an object `x` has shape
80
+ `s`, then the shape of
81
+ `x.mean_times_precision` is `(*s, n)` and the shape of
82
+ `x.negative_half_precision` is `(*s, n, n)`.
81
83
 
82
84
  Parametrizations
83
85
  ----------------
84
86
  Each exponential family distribution has two special parametrizations: the natural and the
85
87
  expectation parametrization. (These are described in the overview pdf.)
86
88
  Consequently, every distribution has at least two base classes, one inheriting from
87
- :python:`NaturalParametrization` and one from :python:`ExpectationParametrization`.
89
+ `NaturalParametrization` and one from `ExpectationParametrization`.
88
90
 
89
91
  The motivation for the natural parametrization is combining and scaling independent predictive
90
92
  evidence. In the natural parametrization, these operations correspond to scaling and addition.
@@ -94,56 +96,121 @@ maximum likelihood distribution that could have produced them. In the expectati
94
96
  this is an expected value.
95
97
 
96
98
  EFAX provides conversions between the two parametrizations through the
97
- :python:`NaturalParametrization.to_exp` and :python:`ExpectationParametrization.to_nat` methods.
99
+ `NaturalParametrization.to_exp` and `ExpectationParametrization.to_nat` methods.
100
+
101
+ Some distributions also provide additional convenience parametrizations. For example, the normal
102
+ distribution offers `NormalVP` (variance parametrization) and `NormalDP` (deviation
103
+ parametrization), and the multivariate normal provides `MultivariateNormalVP` and
104
+ `MultivariateDiagonalNormalVP`. These are not exponential-family parametrizations but are
105
+ useful for constructing and interpreting distributions.
98
106
 
99
107
  Important methods
100
108
  -----------------
101
109
  EFAX aims to provide the main methods used in machine learning.
102
110
 
103
- Every :python:`Distribution` has methods:
111
+ Every `Distribution` has:
104
112
 
105
- - :python:`flattened` and :python:`unflattened` to flatten and unflatten the parameters into a
106
- single array. Typically, array-valued signals in a machine learning model would be unflattened
107
- into a distribution object, operated on, and then flattened before being sent back to the model.
108
- Flattening is careful with distributions with symmetric (or Hermitian) matrix-valued parameters.
109
- It only stores the upper triangular elements. And,
110
- - :python:`shape`, which supports broadcasting.
113
+ - `shape` and `ndim`, which support broadcasting, and
114
+ - indexing via `[]`, which slices all parameter arrays simultaneously.
111
115
 
112
- Every :python:`NaturalParametrization` has methods:
116
+ Every `NaturalParametrization` has methods:
113
117
 
114
- - :python:`to_exp` to convert itself to expectation parameters.
115
- - :python:`sufficient_statistics` to produce the sufficient statistics given an observation (used in
118
+ - `to_exp` to convert itself to expectation parameters,
119
+ - `sufficient_statistics` to produce the sufficient statistics given an observation (used in
116
120
  maximum likelihood estimation),
117
- - :python:`pdf` and :python:`log_pdf`, which is the density or mass function and its logarithm,
118
- - :python:`fisher_information`, which is the Fisher information matrix, and
119
- - :python:`kl_divergence`, which is the KL divergence.
120
-
121
- Every :python:`ExpectationParametrization` has methods:
122
-
123
- - :python:`to_nat` to convert itself to natural parameters, and
124
- - :python:`kl_divergence`, which is the KL divergence.
121
+ - `log_normalizer`, the log partition function,
122
+ - `carrier_measure`, the base measure,
123
+ - `pdf` and `log_pdf`, which are the density or mass function and its logarithm,
124
+ - `fisher_information_diagonal` and `fisher_information_trace`, which return the
125
+ diagonal and trace of the Fisher information matrix stored as distribution objects,
126
+ - `apply_fisher_information`, which applies the Fisher information matrix to a vector of
127
+ expectation parameters efficiently in a single VJP pass,
128
+ - `jeffreys_prior_density`, which returns the square root of the Fisher information
129
+ determinant,
130
+ - `characteristic_function`, which evaluates the characteristic function of the sufficient
131
+ statistics via analytic continuation of the log-normalizer, and
132
+ - `kl_divergence`, which is the KL divergence.
133
+
134
+ Every `ExpectationParametrization` has methods:
135
+
136
+ - `to_nat` to convert itself to natural parameters, and
137
+ - `kl_divergence`, which is the KL divergence.
125
138
 
126
139
  Some parametrizations inherit from these interfaces:
127
140
 
128
- - :python:`HasConjugatePrior` can produce the conjugate prior,
129
- - :python:`HasGeneralizedConjugatePrior` can produce a generalization of the conjugate prior,
130
- - :python:`Multidimensional` distributions have a integer number of `dimensions`, and
131
- - :python:`Samplable` distributions support sampling.
141
+ - `HasConjugatePrior` can produce and recover the conjugate prior,
142
+ - `HasGeneralizedConjugatePrior` extends that with per-dimension pseudo-observation counts,
143
+ - `Multidimensional` distributions have an integer number of `dimensions`, and
144
+ - `Samplable` distributions support sampling.
132
145
 
133
146
  Some parametrizations inherit from these public mixins:
134
147
 
135
- - :python:`HasEntropyEP` is an expectation parametrization with an entropy and cross entropy, and
136
- - :python:`HasEntropyNP` is a natural parametrization with an entropy, (The cross entropy is not
137
- efficient.)
148
+ - `HasEntropy` is a distribution with a `entropy` method,
149
+ - `HasEntropyEP` is an expectation parametrization with analytically tractable entropy and
150
+ `cross_entropy`, and
151
+ - `HasEntropyNP` is a natural parametrization with analytically tractable entropy via the
152
+ paired expectation parametrization.
138
153
 
139
154
  Some parametrizations inherit from these private mixins:
140
155
 
141
- - :python:`ExpToNat` implements the conversion from expectation to natural parameters when no
156
+ - `ExpToNat` implements the conversion from expectation to natural parameters when no
142
157
  analytical solution is possible. It uses Newton's method with a Jacobian to invert the gradient
143
158
  log-normalizer.
144
- - :python:`TransformedNaturalParametrization` produces a natural parametrization by relating it to
159
+ - `TransformedNaturalParametrization` produces a natural parametrization by relating it to
145
160
  an existing natural parametrization. And similarly for
146
- :python:`TransformedExpectationParametrization`.
161
+ `TransformedExpectationParametrization`.
162
+
163
+ Joint distributions
164
+ -------------------
165
+ `JointDistribution`, `JointDistributionE`, and `JointDistributionN` compose
166
+ multiple independent distributions into a single object. `JointDistributionE` holds
167
+ expectation parametrizations and implements `HasEntropyEP`; `JointDistributionN`
168
+ holds natural parametrizations. They support the same `to_nat` / `to_exp`
169
+ conversions as simple distributions.
170
+
171
+ Structure utilities
172
+ -------------------
173
+ EFAX provides three classes that capture the static metadata of a distribution tree—its types,
174
+ parameter names, and dimension information—without requiring a live instance. They form an
175
+ inheritance hierarchy:
176
+
177
+ **Assembler**
178
+ Stores a post-order traversal of a `Distribution` tree (types, paths, dimensions) so
179
+ that distributions can be reconstructed from raw parameter data without passing type information
180
+ alongside arrays. Key methods:
181
+
182
+ - `assemble(params)` — rebuild a `Distribution` from a `{path: array}` mapping,
183
+ - `coerce_from_distribution(q)` — reinterpret q's numeric values under this Assembler's types,
184
+ - `domain_support()` — enumerate each leaf distribution's parameter constraints,
185
+ - `generate_random(xp, rng, shape, safety)` — draw a random distribution with valid parameters, and
186
+ - `to_nat()` / `to_exp()` — return a copy whose types are all in natural or expectation form.
187
+
188
+ **Estimator** *(extends Assembler)*
189
+ Adds maximum likelihood estimation by recording which parameters are fixed (held constant across
190
+ observations) and which are free. Because the MLE for every exponential family equals the mean
191
+ of the sufficient statistics, estimation reduces to a single call:
192
+
193
+ - `sufficient_statistics(x)` — compute the sufficient statistics of observation x, with
194
+ fixed parameters supplied automatically.
195
+
196
+ Create one with `Estimator.from_type(type_p, **fixed)`,
197
+ `Estimator.from_expectation(p)`, or `Estimator.from_natural(p)`.
198
+
199
+ **Flattener** *(extends Estimator)*
200
+ Adds encoding and decoding between a `Distribution` and an array of shape
201
+ `(*distribution.shape, k)`, making distributions compatible with neural networks and numerical
202
+ optimizers. Fixed parameters are excluded from the encoded array and reinserted automatically
203
+ on decode.
204
+
205
+ - `Flattener.flatten(p, mapped_to_plane=True)` — encode p into a `(Flattener, array)` pair,
206
+ - `unflatten(array)` — decode the array back into a distribution, and
207
+ - `final_dimension_size()` — the size k of the last axis of the encoded array.
208
+
209
+ The `mapped_to_plane` flag controls whether constrained parameters (e.g., those on a
210
+ simplex or restricted to the positive reals) are bijectively mapped to all of ℝⁿ. Set it
211
+ `True` when passing to a neural network (to prevent invalid outputs), and `False`
212
+ when the raw magnitudes matter—for example when differencing expectation parameters or computing
213
+ Jacobians.
147
214
 
148
215
  Distributions
149
216
  =============
@@ -187,7 +254,7 @@ EFAX supports the following distributions:
187
254
  - on a finite set:
188
255
 
189
256
  - Bernoulli
190
- - multinomial
257
+ - categorical
191
258
 
192
259
  - on the nonnegative integers:
193
260
 
@@ -216,6 +283,11 @@ EFAX supports the following distributions:
216
283
  - on the n-sphere:
217
284
 
218
285
  - von Mises-Fisher
286
+ - complex von Mises
287
+
288
+ - on positive-definite matrices:
289
+
290
+ - Wishart
219
291
 
220
292
  Usage
221
293
  =====
@@ -404,14 +476,14 @@ instead.
404
476
  likelihood estimation.
405
477
 
406
478
  Suppose you have some samples from a distribution family with unknown
407
- parameters, and you want to estimate the maximum likelihood parmaters of the
479
+ parameters, and you want to estimate the maximum likelihood parameters of the
408
480
  distribution.
409
481
  """
410
482
  import jax.numpy as jnp
411
483
  import jax.random as jr
412
484
  from tjax import print_generic
413
485
 
414
- from efax import DirichletEP, DirichletNP, MaximumLikelihoodEstimator, parameter_mean
486
+ from efax import DirichletEP, DirichletNP, Estimator, parameter_mean
415
487
 
416
488
  # Consider a Dirichlet distribution with a given alpha.
417
489
  alpha = jnp.asarray([2.0, 3.0, 4.0])
@@ -424,9 +496,9 @@ instead.
424
496
 
425
497
  # Now, let's find the maximum likelihood Dirichlet distribution that fits it.
426
498
  # First, convert the samples to their sufficient statistics.
427
- estimator = MaximumLikelihoodEstimator.create_simple_estimator(DirichletEP)
499
+ estimator = Estimator.from_type(DirichletEP)
428
500
  ss = estimator.sufficient_statistics(samples)
429
- # ss has type DirichletEP. This is similar to the conjguate prior of the
501
+ # ss has type DirichletEP. This is similar to the conjugate prior of the
430
502
  # Dirichlet distribution.
431
503
 
432
504
  # Take the mean over the first axis.
@@ -31,7 +31,7 @@ from ._src.distributions.log_normal.unit_variance import (
31
31
  UnitVarianceLogNormalNP,
32
32
  )
33
33
  from ._src.distributions.logarithmic import LogarithmicEP, LogarithmicNP
34
- from ._src.distributions.multinomial import MultinomialEP, MultinomialNP
34
+ from ._src.distributions.categorical import CategoricalEP, CategoricalNP
35
35
  from ._src.distributions.multivariate_normal.arbitrary import (
36
36
  MultivariateNormalEP,
37
37
  MultivariateNormalNP,
@@ -92,13 +92,14 @@ from ._src.parameter.support import (
92
92
  VectorSupport,
93
93
  )
94
94
  from ._src.parametrization import Distribution, SimpleDistribution
95
- from ._src.structure.estimator import MaximumLikelihoodEstimator
95
+ from ._src.structure.assembler import Assembler, SubDistributionInfo
96
+ from ._src.structure.estimator import Estimator
96
97
  from ._src.structure.flattener import Flattener
97
- from ._src.structure.structure import Structure, SubDistributionInfo
98
98
  from ._src.tools import parameter_dot_product, parameter_map, parameter_mean
99
99
  from ._src.transform.joint import JointDistribution, JointDistributionE, JointDistributionN
100
100
 
101
101
  __all__ = [
102
+ "Assembler",
102
103
  "BernoulliEP",
103
104
  "BernoulliNP",
104
105
  "BetaEP",
@@ -123,6 +124,7 @@ __all__ = [
123
124
  "DirichletEP",
124
125
  "DirichletNP",
125
126
  "Distribution",
127
+ "Estimator",
126
128
  "ExpectationParametrization",
127
129
  "ExponentialEP",
128
130
  "ExponentialNP",
@@ -153,10 +155,9 @@ __all__ = [
153
155
  "LogNormalNP",
154
156
  "LogarithmicEP",
155
157
  "LogarithmicNP",
156
- "MaximumLikelihoodEstimator",
157
158
  "Multidimensional",
158
- "MultinomialEP",
159
- "MultinomialNP",
159
+ "CategoricalEP",
160
+ "CategoricalNP",
160
161
  "MultivariateDiagonalNormalEP",
161
162
  "MultivariateDiagonalNormalNP",
162
163
  "MultivariateDiagonalNormalVP",
@@ -187,7 +188,6 @@ __all__ = [
187
188
  "SoftplusNormalEP",
188
189
  "SoftplusNormalNP",
189
190
  "SquareMatrixSupport",
190
- "Structure",
191
191
  "SubDistributionInfo",
192
192
  "Support",
193
193
  "SymmetricMatrixSupport",
@@ -15,6 +15,7 @@ from efax._src.mixins.has_entropy import HasEntropyEP, HasEntropyNP
15
15
  from efax._src.natural_parametrization import NaturalParametrization
16
16
  from efax._src.parameter import RealField, ScalarSupport, boolean_ring, distribution_parameter
17
17
  from efax._src.parametrization import SimpleDistribution
18
+
18
19
  from .beta import BetaNP
19
20
 
20
21
 
@@ -140,5 +141,5 @@ class BernoulliEP(HasEntropyEP[BernoulliNP], HasConjugatePrior, Samplable):
140
141
  return (cls(probability), n)
141
142
 
142
143
  @override
143
- def conjugate_prior_observation(self) -> JaxRealArray:
144
+ def as_conjugate_prior_observation(self) -> JaxRealArray:
144
145
  return self.probability
@@ -10,6 +10,7 @@ from tjax.dataclasses import dataclass
10
10
  from efax._src.expectation_parametrization import ExpectationParametrization
11
11
  from efax._src.natural_parametrization import NaturalParametrization
12
12
  from efax._src.parameter import RealField, ScalarSupport
13
+
13
14
  from .dirichlet_common import DirichletCommonEP, DirichletCommonNP
14
15
 
15
16