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.
- {efax-2.1.2 → efax-2.2.0}/PKG-INFO +122 -50
- {efax-2.1.2 → efax-2.2.0}/README.rst +115 -43
- {efax-2.1.2 → efax-2.2.0}/efax/__init__.py +7 -7
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/bernoulli.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/beta.py +1 -0
- efax-2.1.2/efax/_src/distributions/multinomial.py → efax-2.2.0/efax/_src/distributions/categorical.py +19 -18
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/chi.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/unit_variance.py +1 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/unit_variance.py +1 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/dirichlet.py +1 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/exponential.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/gamma.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/geometric.py +1 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/inverse_gamma.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/log_normal.py +2 -2
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/unit_variance.py +4 -4
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/fixed_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/unit_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/negative_binomial.py +1 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/unit_variance.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/poisson.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/rayleigh.py +2 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/softplus.py +2 -2
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/unit_variance.py +4 -4
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/von_mises.py +3 -3
- {efax-2.1.2 → efax-2.2.0}/efax/_src/expectation_parametrization.py +7 -6
- efax-2.2.0/efax/_src/interfaces/conjugate_prior.py +72 -0
- efax-2.2.0/efax/_src/interfaces/multidimensional.py +18 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/interfaces/samplable.py +9 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +12 -3
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/has_entropy.py +23 -8
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/transformed_parametrization.py +36 -9
- {efax-2.1.2 → efax-2.2.0}/efax/_src/natural_parametrization.py +18 -12
- {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/support.py +1 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/parametrization.py +24 -5
- efax-2.1.2/efax/_src/structure/structure.py → efax-2.2.0/efax/_src/structure/assembler.py +47 -23
- {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/estimator.py +51 -23
- {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/flattener.py +46 -34
- {efax-2.1.2 → efax-2.2.0}/efax/_src/tools.py +3 -3
- {efax-2.1.2 → efax-2.2.0}/pyproject.toml +6 -6
- 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.0}/efax/_src/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/chi_square.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/complex_von_mises.py +1 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/dirichlet_common.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/gen_dirichlet.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/inverse_gaussian.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/log_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/logarithmic.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/negative_binomial_common.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/normal/normal.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/weibull.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/distributions/wishart.py +1 -1
- {efax-2.1.2 → efax-2.2.0}/efax/_src/interfaces/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/iteration.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/mixins/exp_to_nat/optimistix.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/parameter.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/parameter/ring.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/parameter_names.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/structure/parameter_supports.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/transform/__init__.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/transform/joint.py +0 -0
- {efax-2.1.2 → efax-2.2.0}/efax/_src/types.py +0 -0
- {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.
|
|
3
|
+
Version: 2.2.0
|
|
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
|
|
@@ -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:
|
|
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
|
|
93
|
-
a modification of Python's dataclasses to support JAX's
|
|
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
|
|
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
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
144
|
+
Every `Distribution` has:
|
|
137
145
|
|
|
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.
|
|
146
|
+
- `shape` and `ndim`, which support broadcasting, and
|
|
147
|
+
- indexing via `[]`, which slices all parameter arrays simultaneously.
|
|
144
148
|
|
|
145
|
-
Every
|
|
149
|
+
Every `NaturalParametrization` has methods:
|
|
146
150
|
|
|
147
|
-
-
|
|
148
|
-
-
|
|
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
|
-
-
|
|
151
|
-
-
|
|
152
|
-
-
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
-
|
|
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
|
-
-
|
|
162
|
-
-
|
|
163
|
-
-
|
|
164
|
-
-
|
|
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
|
-
-
|
|
169
|
-
-
|
|
170
|
-
|
|
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
|
-
-
|
|
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
|
-
-
|
|
192
|
+
- `TransformedNaturalParametrization` produces a natural parametrization by relating it to
|
|
178
193
|
an existing natural parametrization. And similarly for
|
|
179
|
-
|
|
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
|
-
-
|
|
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
|
|
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,
|
|
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 =
|
|
532
|
+
estimator = Estimator.from_type(DirichletEP)
|
|
461
533
|
ss = estimator.sufficient_statistics(samples)
|
|
462
|
-
# ss has type DirichletEP. This is similar to 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:
|
|
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
|
|
60
|
-
a modification of Python's dataclasses to support JAX's
|
|
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
|
|
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
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
111
|
+
Every `Distribution` has:
|
|
104
112
|
|
|
105
|
-
-
|
|
106
|
-
|
|
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
|
|
116
|
+
Every `NaturalParametrization` has methods:
|
|
113
117
|
|
|
114
|
-
-
|
|
115
|
-
-
|
|
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
|
-
-
|
|
118
|
-
-
|
|
119
|
-
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
-
|
|
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
|
-
-
|
|
129
|
-
-
|
|
130
|
-
-
|
|
131
|
-
-
|
|
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
|
-
-
|
|
136
|
-
-
|
|
137
|
-
|
|
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
|
-
-
|
|
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
|
-
-
|
|
159
|
+
- `TransformedNaturalParametrization` produces a natural parametrization by relating it to
|
|
145
160
|
an existing natural parametrization. And similarly for
|
|
146
|
-
|
|
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
|
-
-
|
|
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
|
|
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,
|
|
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 =
|
|
499
|
+
estimator = Estimator.from_type(DirichletEP)
|
|
428
500
|
ss = estimator.sufficient_statistics(samples)
|
|
429
|
-
# ss has type DirichletEP. This is similar to 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.
|
|
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.
|
|
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
|
-
"
|
|
159
|
-
"
|
|
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
|
|
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
|
|