efax 2.2.4__tar.gz → 2.3.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.2.4 → efax-2.3.0}/PKG-INFO +1 -1
- {efax-2.2.4 → efax-2.3.0}/efax/__init__.py +17 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/categorical.py +2 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/gen_dirichlet.py +119 -6
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/log_normal.py +2 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/unit_variance.py +3 -3
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/fixed_variance.py +4 -1
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/weibull.py +2 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +2 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -2
- {efax-2.2.4 → efax-2.3.0}/efax/_src/natural_parametrization.py +17 -6
- {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/__init__.py +2 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/ring.py +26 -5
- {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/support.py +84 -10
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/assembler.py +67 -21
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/estimator.py +16 -8
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/flattener.py +25 -10
- {efax-2.2.4 → efax-2.3.0}/efax/_src/tools.py +48 -15
- {efax-2.2.4 → efax-2.3.0}/efax/_src/transform/joint.py +1 -2
- {efax-2.2.4 → efax-2.3.0}/pyproject.toml +3 -12
- {efax-2.2.4 → efax-2.3.0}/README.rst +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/bernoulli.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/beta.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/chi.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/chi_square.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_von_mises.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/dirichlet.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/dirichlet_common.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/exponential.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/gamma.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/geometric.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/inverse_gamma.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/inverse_gaussian.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/logarithmic.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/negative_binomial.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/negative_binomial_common.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/normal.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/poisson.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/rayleigh.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/von_mises.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/wishart.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/expectation_parametrization.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/conjugate_prior.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/multidimensional.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/samplable.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/iteration.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/has_entropy.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/transformed_parametrization.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/parameter.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/parametrization.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/parameter_names.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/parameter_supports.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/transform/__init__.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/_src/types.py +0 -0
- {efax-2.2.4 → efax-2.3.0}/efax/py.typed +0 -0
|
@@ -87,15 +87,26 @@ from ._src.parameter.support import (
|
|
|
87
87
|
ScalarSupport,
|
|
88
88
|
SimplexSupport,
|
|
89
89
|
SquareMatrixSupport,
|
|
90
|
+
SubsimplexSupport,
|
|
90
91
|
Support,
|
|
91
92
|
SymmetricMatrixSupport,
|
|
92
93
|
VectorSupport,
|
|
93
94
|
)
|
|
94
95
|
from ._src.parametrization import Distribution, SimpleDistribution
|
|
95
|
-
from ._src.structure.assembler import
|
|
96
|
+
from ._src.structure.assembler import (
|
|
97
|
+
Assembler,
|
|
98
|
+
JointDistributionInfo,
|
|
99
|
+
SimpleDistributionInfo,
|
|
100
|
+
SubDistributionInfo,
|
|
101
|
+
)
|
|
96
102
|
from ._src.structure.estimator import Estimator
|
|
97
103
|
from ._src.structure.flattener import Flattener
|
|
98
|
-
from ._src.tools import
|
|
104
|
+
from ._src.tools import (
|
|
105
|
+
parameter_dot_product,
|
|
106
|
+
parameter_holomorphic_dot,
|
|
107
|
+
parameter_map,
|
|
108
|
+
parameter_mean,
|
|
109
|
+
)
|
|
99
110
|
from ._src.transform.joint import JointDistribution, JointDistributionE, JointDistributionN
|
|
100
111
|
|
|
101
112
|
__all__ = [
|
|
@@ -152,6 +163,7 @@ __all__ = [
|
|
|
152
163
|
"IsotropicNormalNP",
|
|
153
164
|
"JointDistribution",
|
|
154
165
|
"JointDistributionE",
|
|
166
|
+
"JointDistributionInfo",
|
|
155
167
|
"JointDistributionN",
|
|
156
168
|
"LogNormalEP",
|
|
157
169
|
"LogNormalNP",
|
|
@@ -184,11 +196,13 @@ __all__ = [
|
|
|
184
196
|
"Samplable",
|
|
185
197
|
"ScalarSupport",
|
|
186
198
|
"SimpleDistribution",
|
|
199
|
+
"SimpleDistributionInfo",
|
|
187
200
|
"SimplexSupport",
|
|
188
201
|
"SoftplusNormalEP",
|
|
189
202
|
"SoftplusNormalNP",
|
|
190
203
|
"SquareMatrixSupport",
|
|
191
204
|
"SubDistributionInfo",
|
|
205
|
+
"SubsimplexSupport",
|
|
192
206
|
"Support",
|
|
193
207
|
"SymmetricMatrixSupport",
|
|
194
208
|
"UnitVarianceLogNormalEP",
|
|
@@ -210,6 +224,7 @@ __all__ = [
|
|
|
210
224
|
"flat_dict_of_parameters",
|
|
211
225
|
"flatten_mapping",
|
|
212
226
|
"parameter_dot_product",
|
|
227
|
+
"parameter_holomorphic_dot",
|
|
213
228
|
"parameter_map",
|
|
214
229
|
"parameter_mean",
|
|
215
230
|
"parameters",
|
|
@@ -78,9 +78,9 @@ class CategoricalNP(
|
|
|
78
78
|
if shape is not None:
|
|
79
79
|
shape += self.shape
|
|
80
80
|
logits = xp.concat((self.log_odds, xp.zeros((*self.shape, 1))), axis=-1)
|
|
81
|
-
retval = xpx.one_hot(jr.categorical(key, logits, shape=shape), self.dimensions() + 1) #
|
|
81
|
+
retval = xpx.one_hot(jr.categorical(key, logits, shape=shape), self.dimensions() + 1) # ty: ignore
|
|
82
82
|
assert isinstance(retval, JaxArray)
|
|
83
|
-
return retval[..., :-1]
|
|
83
|
+
return retval[..., :-1]
|
|
84
84
|
|
|
85
85
|
@override
|
|
86
86
|
def dimensions(self) -> int:
|
|
@@ -12,14 +12,20 @@ from typing import override
|
|
|
12
12
|
|
|
13
13
|
import jax.scipy.special as jss
|
|
14
14
|
from array_api_compat import array_namespace
|
|
15
|
-
from tjax import JaxArray, JaxRealArray, Shape, softplus
|
|
15
|
+
from tjax import JaxArray, JaxRealArray, Shape, inverse_softplus, jit, softplus
|
|
16
16
|
from tjax.dataclasses import dataclass
|
|
17
17
|
|
|
18
18
|
from efax._src.interfaces.multidimensional import Multidimensional
|
|
19
19
|
from efax._src.mixins.exp_to_nat.exp_to_nat import ExpToNat
|
|
20
20
|
from efax._src.mixins.has_entropy import HasEntropyEP, HasEntropyNP
|
|
21
21
|
from efax._src.natural_parametrization import NaturalParametrization
|
|
22
|
-
from efax._src.parameter import
|
|
22
|
+
from efax._src.parameter import (
|
|
23
|
+
RealField,
|
|
24
|
+
SubsimplexSupport,
|
|
25
|
+
VectorSupport,
|
|
26
|
+
distribution_parameter,
|
|
27
|
+
negative_support,
|
|
28
|
+
)
|
|
23
29
|
|
|
24
30
|
|
|
25
31
|
@dataclass
|
|
@@ -42,8 +48,8 @@ class GeneralizedDirichletNP(
|
|
|
42
48
|
|
|
43
49
|
@override
|
|
44
50
|
@classmethod
|
|
45
|
-
def domain_support(cls) ->
|
|
46
|
-
return
|
|
51
|
+
def domain_support(cls) -> SubsimplexSupport:
|
|
52
|
+
return SubsimplexSupport()
|
|
47
53
|
|
|
48
54
|
@override
|
|
49
55
|
def log_normalizer(self) -> JaxRealArray:
|
|
@@ -119,14 +125,56 @@ class GeneralizedDirichletEP(
|
|
|
119
125
|
|
|
120
126
|
@override
|
|
121
127
|
@classmethod
|
|
122
|
-
def domain_support(cls) ->
|
|
123
|
-
return
|
|
128
|
+
def domain_support(cls) -> SubsimplexSupport:
|
|
129
|
+
return SubsimplexSupport()
|
|
124
130
|
|
|
125
131
|
@classmethod
|
|
126
132
|
@override
|
|
127
133
|
def natural_parametrization_cls(cls) -> type[GeneralizedDirichletNP]:
|
|
128
134
|
return GeneralizedDirichletNP
|
|
129
135
|
|
|
136
|
+
@jit
|
|
137
|
+
@override
|
|
138
|
+
def to_nat(self) -> GeneralizedDirichletNP:
|
|
139
|
+
xp = array_namespace(self)
|
|
140
|
+
zero = xp.zeros_like(self.mean_log_cumulative_probability[..., :1])
|
|
141
|
+
beta_bar = xp.concat(
|
|
142
|
+
[
|
|
143
|
+
self.mean_log_cumulative_probability[..., :1],
|
|
144
|
+
xp.diff(self.mean_log_cumulative_probability, axis=-1),
|
|
145
|
+
],
|
|
146
|
+
axis=-1,
|
|
147
|
+
)
|
|
148
|
+
alpha_bar_direct = self.mean_log_probability - xp.concat(
|
|
149
|
+
[zero, self.mean_log_cumulative_probability[..., :-1]],
|
|
150
|
+
axis=-1,
|
|
151
|
+
)
|
|
152
|
+
alpha, beta = _beta_parameters_from_expected_logs(alpha_bar_direct, beta_bar)
|
|
153
|
+
alpha_roll = xp.concat([alpha[..., 1:], zero], axis=-1)
|
|
154
|
+
gamma = -xp.diff(beta, axis=-1, append=xp.ones_like(zero)) - alpha_roll
|
|
155
|
+
return GeneralizedDirichletNP(alpha - 1.0, gamma)
|
|
156
|
+
|
|
157
|
+
@override
|
|
158
|
+
def initial_search_parameters(self) -> JaxRealArray:
|
|
159
|
+
xp = array_namespace(self)
|
|
160
|
+
beta_bar = xp.concat(
|
|
161
|
+
[
|
|
162
|
+
self.mean_log_cumulative_probability[..., :1],
|
|
163
|
+
xp.diff(self.mean_log_cumulative_probability, axis=-1),
|
|
164
|
+
],
|
|
165
|
+
axis=-1,
|
|
166
|
+
)
|
|
167
|
+
zero = xp.zeros_like(self.mean_log_cumulative_probability[..., :1])
|
|
168
|
+
alpha_bar_direct = self.mean_log_probability - xp.concat(
|
|
169
|
+
[zero, self.mean_log_cumulative_probability[..., :-1]],
|
|
170
|
+
axis=-1,
|
|
171
|
+
)
|
|
172
|
+
alpha, beta = _beta_parameters_from_expected_logs(alpha_bar_direct, beta_bar)
|
|
173
|
+
alpha_roll = xp.concat([alpha[..., 1:], zero], axis=-1)
|
|
174
|
+
gamma = -xp.diff(beta, axis=-1, append=xp.ones_like(zero)) - alpha_roll
|
|
175
|
+
gamma = xp.maximum(gamma, 1e-3)
|
|
176
|
+
return xp.concat([inverse_softplus(alpha), inverse_softplus(gamma)], axis=-1)
|
|
177
|
+
|
|
130
178
|
@override
|
|
131
179
|
def search_to_natural(self, search_parameters: JaxRealArray) -> GeneralizedDirichletNP:
|
|
132
180
|
# Run Newton's method on the whole real hyperspace.
|
|
@@ -144,3 +192,68 @@ class GeneralizedDirichletEP(
|
|
|
144
192
|
@override
|
|
145
193
|
def dimensions(self) -> int:
|
|
146
194
|
return self.mean_log_probability.shape[-1]
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
inverse_digamma_switch = -2.22
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def _inverse_digamma(y: JaxRealArray) -> JaxRealArray:
|
|
201
|
+
xp = array_namespace(y)
|
|
202
|
+
x = xp.where(
|
|
203
|
+
y >= inverse_digamma_switch,
|
|
204
|
+
xp.exp(y) + 0.5,
|
|
205
|
+
-1.0 / (y - jss.digamma(1.0)),
|
|
206
|
+
)
|
|
207
|
+
for _ in range(8):
|
|
208
|
+
x -= (jss.digamma(x) - y) / jss.polygamma(1, x)
|
|
209
|
+
return x
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def _beta_parameters_from_expected_logs(
|
|
213
|
+
mean_log_probability: JaxRealArray,
|
|
214
|
+
mean_log_complement: JaxRealArray,
|
|
215
|
+
) -> tuple[JaxRealArray, JaxRealArray]:
|
|
216
|
+
xp = array_namespace(mean_log_probability, mean_log_complement)
|
|
217
|
+
mean = xp.exp(mean_log_probability)
|
|
218
|
+
complement = xp.exp(mean_log_complement)
|
|
219
|
+
simplex_mean = mean / (mean + complement)
|
|
220
|
+
concentration = xp.maximum(2.0, 1.0 / xp.maximum(simplex_mean * (1.0 - simplex_mean), 1e-3))
|
|
221
|
+
for _ in range(20):
|
|
222
|
+
psi_sum = jss.digamma(concentration)
|
|
223
|
+
alpha = _inverse_digamma(mean_log_probability + psi_sum)
|
|
224
|
+
beta = _inverse_digamma(mean_log_complement + psi_sum)
|
|
225
|
+
concentration = xp.maximum(alpha + beta, 1e-6)
|
|
226
|
+
alpha = xp.maximum(alpha, 1e-6)
|
|
227
|
+
beta = xp.maximum(beta, 1e-6)
|
|
228
|
+
log_alpha = xp.log(alpha)
|
|
229
|
+
log_beta = xp.log(beta)
|
|
230
|
+
for _ in range(20):
|
|
231
|
+
alpha = xp.exp(log_alpha)
|
|
232
|
+
beta = xp.exp(log_beta)
|
|
233
|
+
delta_log_alpha, delta_log_beta = _beta_expected_log_newton_step(
|
|
234
|
+
alpha, beta, mean_log_probability, mean_log_complement
|
|
235
|
+
)
|
|
236
|
+
log_alpha -= delta_log_alpha
|
|
237
|
+
log_beta -= delta_log_beta
|
|
238
|
+
alpha = xp.exp(log_alpha)
|
|
239
|
+
beta = xp.exp(log_beta)
|
|
240
|
+
return alpha, beta
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def _beta_expected_log_newton_step(
|
|
244
|
+
alpha: JaxRealArray,
|
|
245
|
+
beta: JaxRealArray,
|
|
246
|
+
mean_log_probability: JaxRealArray,
|
|
247
|
+
mean_log_complement: JaxRealArray,
|
|
248
|
+
) -> tuple[JaxRealArray, JaxRealArray]:
|
|
249
|
+
trigamma_sum = jss.polygamma(1, alpha + beta)
|
|
250
|
+
f1 = jss.digamma(alpha) - jss.digamma(alpha + beta) - mean_log_probability
|
|
251
|
+
f2 = jss.digamma(beta) - jss.digamma(alpha + beta) - mean_log_complement
|
|
252
|
+
j11 = alpha * (jss.polygamma(1, alpha) - trigamma_sum)
|
|
253
|
+
j12 = -beta * trigamma_sum
|
|
254
|
+
j21 = -alpha * trigamma_sum
|
|
255
|
+
j22 = beta * (jss.polygamma(1, beta) - trigamma_sum)
|
|
256
|
+
determinant = j11 * j22 - j12 * j21
|
|
257
|
+
delta_log_alpha = (j22 * f1 - j12 * f2) / determinant
|
|
258
|
+
delta_log_beta = (-j21 * f1 + j11 * f2) / determinant
|
|
259
|
+
return delta_log_alpha, delta_log_beta
|
|
@@ -42,7 +42,7 @@ class LogNormalNP(
|
|
|
42
42
|
@override
|
|
43
43
|
@classmethod
|
|
44
44
|
def domain_support(cls) -> ScalarSupport:
|
|
45
|
-
return ScalarSupport()
|
|
45
|
+
return ScalarSupport(ring=positive_support)
|
|
46
46
|
|
|
47
47
|
@override
|
|
48
48
|
@classmethod
|
|
@@ -103,7 +103,7 @@ class LogNormalEP(
|
|
|
103
103
|
@override
|
|
104
104
|
@classmethod
|
|
105
105
|
def domain_support(cls) -> ScalarSupport:
|
|
106
|
-
return ScalarSupport()
|
|
106
|
+
return ScalarSupport(ring=positive_support)
|
|
107
107
|
|
|
108
108
|
@classmethod
|
|
109
109
|
@override
|
|
@@ -17,7 +17,7 @@ from efax._src.mixins.transformed_parametrization import (
|
|
|
17
17
|
TransformedNaturalParametrization,
|
|
18
18
|
)
|
|
19
19
|
from efax._src.natural_parametrization import NaturalParametrization
|
|
20
|
-
from efax._src.parameter import ScalarSupport, distribution_parameter
|
|
20
|
+
from efax._src.parameter import ScalarSupport, distribution_parameter, positive_support
|
|
21
21
|
|
|
22
22
|
|
|
23
23
|
@dataclass
|
|
@@ -35,7 +35,7 @@ class UnitVarianceLogNormalNP(
|
|
|
35
35
|
@override
|
|
36
36
|
@classmethod
|
|
37
37
|
def domain_support(cls) -> ScalarSupport:
|
|
38
|
-
return ScalarSupport()
|
|
38
|
+
return ScalarSupport(ring=positive_support)
|
|
39
39
|
|
|
40
40
|
@override
|
|
41
41
|
@classmethod
|
|
@@ -90,7 +90,7 @@ class UnitVarianceLogNormalEP(
|
|
|
90
90
|
@override
|
|
91
91
|
@classmethod
|
|
92
92
|
def domain_support(cls) -> ScalarSupport:
|
|
93
|
-
return ScalarSupport()
|
|
93
|
+
return ScalarSupport(ring=positive_support)
|
|
94
94
|
|
|
95
95
|
@classmethod
|
|
96
96
|
@override
|
|
@@ -76,7 +76,10 @@ class MultivariateFixedVarianceNormalNP(
|
|
|
76
76
|
def sufficient_statistics(
|
|
77
77
|
cls, x: JaxRealArray, **fixed_parameters: JaxArray
|
|
78
78
|
) -> MultivariateFixedVarianceNormalEP:
|
|
79
|
-
|
|
79
|
+
xp = array_namespace(x)
|
|
80
|
+
return MultivariateFixedVarianceNormalEP(
|
|
81
|
+
x, variance=xp.broadcast_to(fixed_parameters["variance"], x.shape[:-1])
|
|
82
|
+
)
|
|
80
83
|
|
|
81
84
|
@override
|
|
82
85
|
def sample(self, key: KeyArray, shape: Shape | None = None) -> JaxRealArray:
|
|
@@ -68,8 +68,8 @@ class WeibullNP(
|
|
|
68
68
|
@classmethod
|
|
69
69
|
def sufficient_statistics(cls, x: JaxRealArray, **fixed_parameters: JaxArray) -> WeibullEP:
|
|
70
70
|
xp = array_namespace(x)
|
|
71
|
-
concentration = fixed_parameters["concentration"]
|
|
72
|
-
return WeibullEP(
|
|
71
|
+
concentration = xp.broadcast_to(fixed_parameters["concentration"], x.shape)
|
|
72
|
+
return WeibullEP(concentration, x**concentration)
|
|
73
73
|
|
|
74
74
|
@override
|
|
75
75
|
def sample(self, key: KeyArray, shape: Shape | None = None) -> JaxRealArray:
|
|
@@ -43,7 +43,7 @@ class ExpToNat(ExpectationParametrization[NP], SimpleDistribution, Generic[NP]):
|
|
|
43
43
|
# Select a default minimizer based on the dimensionality of the search space.
|
|
44
44
|
# Bisection is used for scalar problems; Newton's method for vector problems.
|
|
45
45
|
if hasattr(super(), "__post_init__"):
|
|
46
|
-
super().__post_init__() #
|
|
46
|
+
super().__post_init__() # ty: ignore
|
|
47
47
|
if self.minimizer is None:
|
|
48
48
|
from .optimistix import default_bisection_minimizer, default_minimizer # noqa: PLC0415
|
|
49
49
|
|
|
@@ -96,7 +96,7 @@ class ExpToNat(ExpectationParametrization[NP], SimpleDistribution, Generic[NP]):
|
|
|
96
96
|
unflatten_as_type=np_cls,
|
|
97
97
|
mapped_to_plane=True,
|
|
98
98
|
)
|
|
99
|
-
return flattener.unflatten(search_parameters) #
|
|
99
|
+
return flattener.unflatten(search_parameters) # ty: ignore
|
|
100
100
|
|
|
101
101
|
def search_gradient(self, search_parameters: SP) -> SP:
|
|
102
102
|
"""Convert the search parameters to the natural gradient.
|
|
@@ -55,13 +55,13 @@ class OptimistixRootFinder(ExpToNatMinimizer):
|
|
|
55
55
|
|
|
56
56
|
|
|
57
57
|
default_minimizer = OptimistixRootFinder(
|
|
58
|
-
solver=optx.Newton[JaxRealArray, JaxRealArray, None](rtol=0.0, atol=1e-7),
|
|
58
|
+
solver=optx.Newton[JaxRealArray, JaxRealArray, None](rtol=0.0, atol=1e-7),
|
|
59
59
|
max_steps=1000,
|
|
60
60
|
)
|
|
61
61
|
|
|
62
62
|
|
|
63
63
|
default_bisection_minimizer = OptimistixRootFinder(
|
|
64
|
-
solver=optx.Bisection( #
|
|
64
|
+
solver=optx.Bisection( # ty: ignore
|
|
65
65
|
rtol=0.0, atol=1e-7, flip="detect", expand_if_necessary=True
|
|
66
66
|
),
|
|
67
67
|
max_steps=1000,
|
|
@@ -24,7 +24,7 @@ from .parametrization import Distribution
|
|
|
24
24
|
from .structure.assembler import Assembler
|
|
25
25
|
from .structure.estimator import Estimator
|
|
26
26
|
from .structure.flattener import Flattener
|
|
27
|
-
from .tools import parameter_dot_product
|
|
27
|
+
from .tools import parameter_dot_product, parameter_holomorphic_dot
|
|
28
28
|
|
|
29
29
|
if TYPE_CHECKING:
|
|
30
30
|
from .expectation_parametrization import ExpectationParametrization
|
|
@@ -37,13 +37,24 @@ Domain = TypeVar("Domain", bound=JaxComplexArray | dict[str, Any], default=Any)
|
|
|
37
37
|
def _log_normalizer_jvp(
|
|
38
38
|
primals: tuple[NaturalParametrization],
|
|
39
39
|
tangents: tuple[NaturalParametrization],
|
|
40
|
-
) -> tuple[
|
|
41
|
-
"""The log-normalizer's special JVP vastly improves numerical stability.
|
|
40
|
+
) -> tuple[JaxArray, JaxArray]:
|
|
41
|
+
"""The log-normalizer's special JVP vastly improves numerical stability.
|
|
42
|
+
|
|
43
|
+
On the analytic-continuation path used by ``characteristic_function`` the
|
|
44
|
+
natural parameters live in ℂⁿ and the primal output ``y`` is complex; the
|
|
45
|
+
holomorphic derivative ``Σ q_dot · p`` is what JAX expects as the complex
|
|
46
|
+
tangent. In every other case (real distributions, and intrinsically-complex
|
|
47
|
+
distributions whose ``log_normalizer`` returns a real scalar) the Hermitian
|
|
48
|
+
form is correct. Dispatching on ``y.dtype`` distinguishes the two.
|
|
49
|
+
"""
|
|
42
50
|
(q,) = primals
|
|
43
51
|
(q_dot,) = tangents
|
|
44
52
|
y = q.log_normalizer()
|
|
45
53
|
p = q.to_exp()
|
|
46
|
-
|
|
54
|
+
if jnp.iscomplexobj(y):
|
|
55
|
+
y_dot = parameter_holomorphic_dot(q_dot, p)
|
|
56
|
+
else:
|
|
57
|
+
y_dot = parameter_dot_product(q_dot, p)
|
|
47
58
|
return y, y_dot
|
|
48
59
|
|
|
49
60
|
|
|
@@ -157,7 +168,7 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
|
|
|
157
168
|
"""
|
|
158
169
|
xp = array_namespace(self)
|
|
159
170
|
fisher_information_diagonal = self.fisher_information_diagonal()
|
|
160
|
-
|
|
171
|
+
assembler = Assembler.create_assembler(self)
|
|
161
172
|
final_parameters = parameters(self)
|
|
162
173
|
for path, (value, support) in parameters(fisher_information_diagonal, support=True).items():
|
|
163
174
|
na = support.axes()
|
|
@@ -170,7 +181,7 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
|
|
|
170
181
|
else:
|
|
171
182
|
raise RuntimeError
|
|
172
183
|
final_parameters[path] = new_value
|
|
173
|
-
return
|
|
184
|
+
return assembler.assemble(final_parameters)
|
|
174
185
|
|
|
175
186
|
@final
|
|
176
187
|
def jeffreys_prior_density(self) -> JaxRealArray:
|
|
@@ -17,6 +17,7 @@ from .support import (
|
|
|
17
17
|
ScalarSupport,
|
|
18
18
|
SimplexSupport,
|
|
19
19
|
SquareMatrixSupport,
|
|
20
|
+
SubsimplexSupport,
|
|
20
21
|
Support,
|
|
21
22
|
SymmetricMatrixSupport,
|
|
22
23
|
VectorSupport,
|
|
@@ -32,6 +33,7 @@ __all__ = [
|
|
|
32
33
|
"ScalarSupport",
|
|
33
34
|
"SimplexSupport",
|
|
34
35
|
"SquareMatrixSupport",
|
|
36
|
+
"SubsimplexSupport",
|
|
35
37
|
"Support",
|
|
36
38
|
"SymmetricMatrixSupport",
|
|
37
39
|
"VectorSupport",
|
|
@@ -16,6 +16,7 @@ from tjax import (
|
|
|
16
16
|
JaxRealArray,
|
|
17
17
|
RealNumeric,
|
|
18
18
|
Shape,
|
|
19
|
+
divide_where,
|
|
19
20
|
inverse_softplus,
|
|
20
21
|
softplus,
|
|
21
22
|
)
|
|
@@ -142,7 +143,7 @@ class RealField(Ring):
|
|
|
142
143
|
minimum = xp.asarray(self.minimum)
|
|
143
144
|
maximum = xp.asarray(self.maximum)
|
|
144
145
|
domain_width = xp.maximum(maximum - minimum - safety, 0.0)
|
|
145
|
-
true_minimum =
|
|
146
|
+
true_minimum = (minimum + maximum - domain_width) * 0.5
|
|
146
147
|
return xp.asarray(true_minimum + domain_width * rng.uniform(size=shape))
|
|
147
148
|
return xp.asarray(rng.normal(scale=1.0, size=shape) * self.generation_scale)
|
|
148
149
|
|
|
@@ -176,12 +177,22 @@ class ComplexField(Ring):
|
|
|
176
177
|
# x is outside the disk of the given minimum. Map it to the plane.
|
|
177
178
|
magnitude = xp.abs(x)
|
|
178
179
|
corrected_magnitude = magnitude - minimum
|
|
179
|
-
corrected_x = x *
|
|
180
|
+
corrected_x = x * divide_where(
|
|
181
|
+
corrected_magnitude,
|
|
182
|
+
magnitude,
|
|
183
|
+
where=magnitude != 0.0,
|
|
184
|
+
otherwise=xp.zeros_like(magnitude),
|
|
185
|
+
)
|
|
180
186
|
else:
|
|
181
187
|
# x is in the disk of the given maximum. Map it to the plane.
|
|
182
188
|
magnitude = xp.abs(x)
|
|
183
189
|
corrected_magnitude = jss.logit(((magnitude - minimum) / maximum) * 0.5 + 0.5)
|
|
184
|
-
corrected_x = x *
|
|
190
|
+
corrected_x = x * divide_where(
|
|
191
|
+
corrected_magnitude,
|
|
192
|
+
magnitude,
|
|
193
|
+
where=magnitude != 0.0,
|
|
194
|
+
otherwise=xp.zeros_like(magnitude),
|
|
195
|
+
)
|
|
185
196
|
else:
|
|
186
197
|
corrected_x = x
|
|
187
198
|
return xp.concat([xp.real(corrected_x), xp.imag(corrected_x)], axis=-1)
|
|
@@ -206,11 +217,21 @@ class ComplexField(Ring):
|
|
|
206
217
|
# x is outside the disk of the given minimum. Map it to the plane.
|
|
207
218
|
corrected_magnitude = xp.abs(corrected_x)
|
|
208
219
|
magnitude = corrected_magnitude + minimum
|
|
209
|
-
return corrected_x *
|
|
220
|
+
return corrected_x * divide_where(
|
|
221
|
+
magnitude,
|
|
222
|
+
corrected_magnitude,
|
|
223
|
+
where=corrected_magnitude != 0.0,
|
|
224
|
+
otherwise=xp.zeros_like(corrected_magnitude),
|
|
225
|
+
)
|
|
210
226
|
# x is in the disk of the given maximum. Map it to the plane.
|
|
211
227
|
corrected_magnitude = xp.abs(corrected_x)
|
|
212
228
|
magnitude = maximum * (jss.expit(corrected_magnitude) - 0.5) * 2.0 + minimum
|
|
213
|
-
return corrected_x *
|
|
229
|
+
return corrected_x * divide_where(
|
|
230
|
+
magnitude,
|
|
231
|
+
corrected_magnitude,
|
|
232
|
+
where=corrected_magnitude != 0.0,
|
|
233
|
+
otherwise=xp.zeros_like(corrected_magnitude),
|
|
234
|
+
)
|
|
214
235
|
|
|
215
236
|
@override
|
|
216
237
|
def generate(self, xp: Namespace, rng: Generator, shape: Shape, safety: float) -> JaxRealArray:
|
|
@@ -10,7 +10,7 @@ import numpy as np
|
|
|
10
10
|
from array_api_compat import array_namespace
|
|
11
11
|
from numpy.random import Generator
|
|
12
12
|
from opt_einsum import contract
|
|
13
|
-
from tjax import JaxArray, JaxRealArray, Shape
|
|
13
|
+
from tjax import JaxArray, JaxRealArray, Shape, divide_where
|
|
14
14
|
|
|
15
15
|
from efax._src.types import Namespace
|
|
16
16
|
|
|
@@ -141,13 +141,22 @@ class SimplexSupport(Support):
|
|
|
141
141
|
|
|
142
142
|
@override
|
|
143
143
|
def flattened(self, x: JaxArray, *, map_to_plane: bool) -> JaxRealArray:
|
|
144
|
-
|
|
144
|
+
if not map_to_plane:
|
|
145
|
+
return self.ring.flattened(x, map_to_plane=False)
|
|
146
|
+
xp = array_namespace(x)
|
|
147
|
+
residual = 1.0 - xp.sum(x, axis=-1, keepdims=True)
|
|
148
|
+
return xp.log(x / residual)
|
|
145
149
|
|
|
146
150
|
@override
|
|
147
151
|
def unflattened(self, y: JaxRealArray, dimensions: int, *, map_from_plane: bool) -> JaxArray:
|
|
148
|
-
|
|
149
|
-
|
|
150
|
-
|
|
152
|
+
if not map_from_plane:
|
|
153
|
+
x = self.ring.unflattened(y, map_from_plane=False)
|
|
154
|
+
assert x.shape[-1] == dimensions - 1
|
|
155
|
+
return x
|
|
156
|
+
xp = array_namespace(y)
|
|
157
|
+
assert y.shape[-1] == dimensions - 1
|
|
158
|
+
logits = xp.concat((y, xp.zeros((*y.shape[:-1], 1), dtype=y.dtype)), axis=-1)
|
|
159
|
+
return jss.softmax(logits, axis=-1)[..., :-1]
|
|
151
160
|
|
|
152
161
|
@override
|
|
153
162
|
def clamp(self, x: JaxArray) -> JaxArray:
|
|
@@ -161,7 +170,62 @@ class SimplexSupport(Support):
|
|
|
161
170
|
def generate(
|
|
162
171
|
self, xp: Namespace, rng: Generator, shape: Shape, safety: float, dimensions: int
|
|
163
172
|
) -> JaxRealArray:
|
|
164
|
-
|
|
173
|
+
alpha = xp.ones(dimensions)
|
|
174
|
+
x = xp.asarray(rng.dirichlet(alpha, size=shape))
|
|
175
|
+
if safety > 0.0:
|
|
176
|
+
x = x * (1.0 - safety) + safety / dimensions
|
|
177
|
+
return x[..., :-1]
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
class SubsimplexSupport(Support):
|
|
181
|
+
@override
|
|
182
|
+
def axes(self) -> int:
|
|
183
|
+
return 1
|
|
184
|
+
|
|
185
|
+
@override
|
|
186
|
+
def shape(self, dimensions: int) -> Shape:
|
|
187
|
+
return (dimensions,)
|
|
188
|
+
|
|
189
|
+
@override
|
|
190
|
+
def num_elements(self, dimensions: int) -> int:
|
|
191
|
+
return self.ring.num_elements(dimensions)
|
|
192
|
+
|
|
193
|
+
@override
|
|
194
|
+
def flattened(self, x: JaxArray, *, map_to_plane: bool) -> JaxRealArray:
|
|
195
|
+
if not map_to_plane:
|
|
196
|
+
return self.ring.flattened(x, map_to_plane=False)
|
|
197
|
+
xp = array_namespace(x)
|
|
198
|
+
residual = 1.0 - xp.sum(x, axis=-1, keepdims=True)
|
|
199
|
+
return xp.log(x / residual)
|
|
200
|
+
|
|
201
|
+
@override
|
|
202
|
+
def unflattened(self, y: JaxRealArray, dimensions: int, *, map_from_plane: bool) -> JaxArray:
|
|
203
|
+
if not map_from_plane:
|
|
204
|
+
x = self.ring.unflattened(y, map_from_plane=False)
|
|
205
|
+
assert x.shape[-1] == dimensions
|
|
206
|
+
return x
|
|
207
|
+
xp = array_namespace(y)
|
|
208
|
+
assert y.shape[-1] == dimensions
|
|
209
|
+
logits = xp.concat((y, xp.zeros((*y.shape[:-1], 1), dtype=y.dtype)), axis=-1)
|
|
210
|
+
return jss.softmax(logits, axis=-1)[..., :-1]
|
|
211
|
+
|
|
212
|
+
@override
|
|
213
|
+
def clamp(self, x: JaxArray) -> JaxArray:
|
|
214
|
+
xp = array_namespace(x)
|
|
215
|
+
eps = xp.finfo(x.dtype).eps
|
|
216
|
+
s = xp.sum(x, axis=-1, keepdims=True)
|
|
217
|
+
x *= xp.minimum(1.0, (1.0 - eps) / s)
|
|
218
|
+
return xp.clip(x, min=eps, max=1.0 - eps)
|
|
219
|
+
|
|
220
|
+
@override
|
|
221
|
+
def generate(
|
|
222
|
+
self, xp: Namespace, rng: Generator, shape: Shape, safety: float, dimensions: int
|
|
223
|
+
) -> JaxRealArray:
|
|
224
|
+
alpha = xp.ones(dimensions + 1)
|
|
225
|
+
x = xp.asarray(rng.dirichlet(alpha, size=shape))[..., :-1]
|
|
226
|
+
if safety > 0.0:
|
|
227
|
+
x *= 1.0 - safety
|
|
228
|
+
return x
|
|
165
229
|
|
|
166
230
|
|
|
167
231
|
class SymmetricMatrixSupport(Support):
|
|
@@ -224,10 +288,10 @@ class SymmetricMatrixSupport(Support):
|
|
|
224
288
|
i = int(i_)
|
|
225
289
|
j = int(j_)
|
|
226
290
|
xk = x[..., k]
|
|
227
|
-
result = xpx.at(result)[..., i, j].set(xk) #
|
|
291
|
+
result = xpx.at(result)[..., i, j].set(xk) # ty: ignore
|
|
228
292
|
if i != j:
|
|
229
293
|
cxk = xp.conj(xk) if self.hermitian else xk
|
|
230
|
-
result = xpx.at(result)[..., j, i].set(cxk) #
|
|
294
|
+
result = xpx.at(result)[..., j, i].set(cxk) # ty: ignore
|
|
231
295
|
assert isinstance(result, JaxArray)
|
|
232
296
|
return result
|
|
233
297
|
|
|
@@ -299,7 +363,12 @@ class CircularBoundedSupport(VectorSupport):
|
|
|
299
363
|
xp = array_namespace(x)
|
|
300
364
|
magnitude = xp.linalg.norm(x, 2, axis=-1, keepdims=True)
|
|
301
365
|
corrected_magnitude = jss.logit((magnitude / self.radius) * 0.5 + 0.5)
|
|
302
|
-
return x *
|
|
366
|
+
return x * divide_where(
|
|
367
|
+
corrected_magnitude,
|
|
368
|
+
magnitude,
|
|
369
|
+
where=magnitude != 0.0,
|
|
370
|
+
otherwise=xp.zeros_like(magnitude),
|
|
371
|
+
)
|
|
303
372
|
|
|
304
373
|
@override
|
|
305
374
|
def unflattened(self, y: JaxRealArray, dimensions: int, *, map_from_plane: bool) -> JaxArray:
|
|
@@ -310,4 +379,9 @@ class CircularBoundedSupport(VectorSupport):
|
|
|
310
379
|
assert y.shape[-1] == dimensions
|
|
311
380
|
corrected_magnitude = cast("JaxRealArray", xp.linalg.norm(y, 2, axis=-1, keepdims=True))
|
|
312
381
|
magnitude = self.radius * (jss.expit(corrected_magnitude) - 0.5) * 2.0
|
|
313
|
-
return y *
|
|
382
|
+
return y * divide_where(
|
|
383
|
+
magnitude,
|
|
384
|
+
corrected_magnitude,
|
|
385
|
+
where=corrected_magnitude != 0.0,
|
|
386
|
+
otherwise=xp.zeros_like(corrected_magnitude),
|
|
387
|
+
)
|