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.
Files changed (78) hide show
  1. {efax-2.2.4 → efax-2.3.0}/PKG-INFO +1 -1
  2. {efax-2.2.4 → efax-2.3.0}/efax/__init__.py +17 -2
  3. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/categorical.py +2 -2
  4. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/gen_dirichlet.py +119 -6
  5. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/log_normal.py +2 -2
  6. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/unit_variance.py +3 -3
  7. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/fixed_variance.py +4 -1
  8. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/weibull.py +2 -2
  9. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +2 -2
  10. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -2
  11. {efax-2.2.4 → efax-2.3.0}/efax/_src/natural_parametrization.py +17 -6
  12. {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/__init__.py +2 -0
  13. {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/ring.py +26 -5
  14. {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/support.py +84 -10
  15. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/assembler.py +67 -21
  16. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/estimator.py +16 -8
  17. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/flattener.py +25 -10
  18. {efax-2.2.4 → efax-2.3.0}/efax/_src/tools.py +48 -15
  19. {efax-2.2.4 → efax-2.3.0}/efax/_src/transform/joint.py +1 -2
  20. {efax-2.2.4 → efax-2.3.0}/pyproject.toml +3 -12
  21. {efax-2.2.4 → efax-2.3.0}/README.rst +0 -0
  22. {efax-2.2.4 → efax-2.3.0}/efax/_src/__init__.py +0 -0
  23. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/__init__.py +0 -0
  24. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/bernoulli.py +0 -0
  25. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/beta.py +0 -0
  26. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/chi.py +0 -0
  27. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/chi_square.py +0 -0
  28. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/__init__.py +0 -0
  29. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
  30. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
  31. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/__init__.py +0 -0
  32. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
  33. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
  34. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/complex_von_mises.py +0 -0
  35. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/dirichlet.py +0 -0
  36. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/dirichlet_common.py +0 -0
  37. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/exponential.py +0 -0
  38. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/gamma.py +0 -0
  39. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/geometric.py +0 -0
  40. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/inverse_gamma.py +0 -0
  41. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/inverse_gaussian.py +0 -0
  42. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/log_normal/__init__.py +0 -0
  43. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/logarithmic.py +0 -0
  44. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
  45. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
  46. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
  47. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
  48. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
  49. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/negative_binomial.py +0 -0
  50. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/negative_binomial_common.py +0 -0
  51. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/__init__.py +0 -0
  52. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/normal.py +0 -0
  53. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/normal/unit_variance.py +0 -0
  54. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/poisson.py +0 -0
  55. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/rayleigh.py +0 -0
  56. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
  57. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
  58. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
  59. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/von_mises.py +0 -0
  60. {efax-2.2.4 → efax-2.3.0}/efax/_src/distributions/wishart.py +0 -0
  61. {efax-2.2.4 → efax-2.3.0}/efax/_src/expectation_parametrization.py +0 -0
  62. {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/__init__.py +0 -0
  63. {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/conjugate_prior.py +0 -0
  64. {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/multidimensional.py +0 -0
  65. {efax-2.2.4 → efax-2.3.0}/efax/_src/interfaces/samplable.py +0 -0
  66. {efax-2.2.4 → efax-2.3.0}/efax/_src/iteration.py +0 -0
  67. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/__init__.py +0 -0
  68. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
  69. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/has_entropy.py +0 -0
  70. {efax-2.2.4 → efax-2.3.0}/efax/_src/mixins/transformed_parametrization.py +0 -0
  71. {efax-2.2.4 → efax-2.3.0}/efax/_src/parameter/parameter.py +0 -0
  72. {efax-2.2.4 → efax-2.3.0}/efax/_src/parametrization.py +0 -0
  73. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/__init__.py +0 -0
  74. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/parameter_names.py +0 -0
  75. {efax-2.2.4 → efax-2.3.0}/efax/_src/structure/parameter_supports.py +0 -0
  76. {efax-2.2.4 → efax-2.3.0}/efax/_src/transform/__init__.py +0 -0
  77. {efax-2.2.4 → efax-2.3.0}/efax/_src/types.py +0 -0
  78. {efax-2.2.4 → efax-2.3.0}/efax/py.typed +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: efax
3
- Version: 2.2.4
3
+ Version: 2.3.0
4
4
  Summary: Exponential families for JAX
5
5
  Author: Neil Girdhar
6
6
  Author-email: Neil Girdhar <mistersheik@gmail.com>
@@ -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 Assembler, SubDistributionInfo
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 parameter_dot_product, parameter_map, parameter_mean
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) # type: ignore
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] # pyright: ignore
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 RealField, VectorSupport, distribution_parameter, negative_support
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) -> VectorSupport:
46
- return VectorSupport()
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) -> VectorSupport:
123
- return VectorSupport()
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
- return MultivariateFixedVarianceNormalEP(x, variance=fixed_parameters["variance"])
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(xp.broadcast_to(concentration, x.shape), x**concentration)
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__() # pyright: ignore
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) # type: ignore
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), # ty: ignore
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( # type: ignore
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[JaxRealArray, JaxRealArray]:
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
- y_dot = parameter_dot_product(q_dot, p)
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
- structure = Assembler.create_assembler(self)
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 structure.assemble(final_parameters)
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 = xp.mean(minimum, maximum) - domain_width * 0.5
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 * corrected_magnitude / magnitude
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 * corrected_magnitude / magnitude
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 * magnitude / corrected_magnitude
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 * magnitude / corrected_magnitude
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
- return self.ring.flattened(x, map_to_plane=map_to_plane)
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
- x = self.ring.unflattened(y, map_from_plane=map_from_plane)
149
- assert x.shape[-1] == dimensions
150
- return x
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
- raise NotImplementedError
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) # type: ignore
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) # type: ignore
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 * corrected_magnitude / magnitude
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 * magnitude / corrected_magnitude
382
+ return y * divide_where(
383
+ magnitude,
384
+ corrected_magnitude,
385
+ where=corrected_magnitude != 0.0,
386
+ otherwise=xp.zeros_like(corrected_magnitude),
387
+ )