efax 2.2.0__tar.gz → 2.2.4__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.0 → efax-2.2.4}/PKG-INFO +36 -11
  2. {efax-2.2.0 → efax-2.2.4}/README.rst +34 -9
  3. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/categorical.py +2 -1
  4. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/gamma.py +7 -4
  5. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/optimistix.py +2 -1
  6. {efax-2.2.0 → efax-2.2.4}/efax/_src/natural_parametrization.py +25 -6
  7. {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/support.py +6 -1
  8. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/assembler.py +7 -3
  9. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/estimator.py +6 -2
  10. {efax-2.2.0 → efax-2.2.4}/efax/_src/tools.py +7 -10
  11. {efax-2.2.0 → efax-2.2.4}/efax/_src/transform/joint.py +13 -3
  12. {efax-2.2.0 → efax-2.2.4}/pyproject.toml +3 -4
  13. {efax-2.2.0 → efax-2.2.4}/efax/__init__.py +3 -3
  14. {efax-2.2.0 → efax-2.2.4}/efax/_src/__init__.py +0 -0
  15. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/__init__.py +0 -0
  16. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/bernoulli.py +0 -0
  17. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/beta.py +0 -0
  18. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/chi.py +0 -0
  19. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/chi_square.py +0 -0
  20. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/__init__.py +0 -0
  21. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
  22. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
  23. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/__init__.py +0 -0
  24. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
  25. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
  26. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/complex_von_mises.py +0 -0
  27. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/dirichlet.py +0 -0
  28. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/dirichlet_common.py +0 -0
  29. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/exponential.py +0 -0
  30. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/gen_dirichlet.py +0 -0
  31. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/geometric.py +0 -0
  32. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/inverse_gamma.py +0 -0
  33. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/inverse_gaussian.py +0 -0
  34. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/__init__.py +0 -0
  35. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/log_normal.py +0 -0
  36. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/log_normal/unit_variance.py +0 -0
  37. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/logarithmic.py +0 -0
  38. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
  39. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
  40. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
  41. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/fixed_variance.py +0 -0
  42. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
  43. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
  44. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/negative_binomial.py +0 -0
  45. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/negative_binomial_common.py +0 -0
  46. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/__init__.py +0 -0
  47. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/normal.py +0 -0
  48. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/normal/unit_variance.py +0 -0
  49. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/poisson.py +0 -0
  50. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/rayleigh.py +0 -0
  51. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
  52. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
  53. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
  54. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/von_mises.py +0 -0
  55. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/weibull.py +0 -0
  56. {efax-2.2.0 → efax-2.2.4}/efax/_src/distributions/wishart.py +0 -0
  57. {efax-2.2.0 → efax-2.2.4}/efax/_src/expectation_parametrization.py +0 -0
  58. {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/__init__.py +0 -0
  59. {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/conjugate_prior.py +0 -0
  60. {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/multidimensional.py +0 -0
  61. {efax-2.2.0 → efax-2.2.4}/efax/_src/interfaces/samplable.py +0 -0
  62. {efax-2.2.0 → efax-2.2.4}/efax/_src/iteration.py +0 -0
  63. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/__init__.py +0 -0
  64. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
  65. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +0 -0
  66. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/has_entropy.py +0 -0
  67. {efax-2.2.0 → efax-2.2.4}/efax/_src/mixins/transformed_parametrization.py +0 -0
  68. {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/__init__.py +0 -0
  69. {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/parameter.py +0 -0
  70. {efax-2.2.0 → efax-2.2.4}/efax/_src/parameter/ring.py +0 -0
  71. {efax-2.2.0 → efax-2.2.4}/efax/_src/parametrization.py +0 -0
  72. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/__init__.py +0 -0
  73. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/flattener.py +0 -0
  74. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/parameter_names.py +0 -0
  75. {efax-2.2.0 → efax-2.2.4}/efax/_src/structure/parameter_supports.py +0 -0
  76. {efax-2.2.0 → efax-2.2.4}/efax/_src/transform/__init__.py +0 -0
  77. {efax-2.2.0 → efax-2.2.4}/efax/_src/types.py +0 -0
  78. {efax-2.2.0 → efax-2.2.4}/efax/py.typed +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: efax
3
- Version: 2.2.0
3
+ Version: 2.2.4
4
4
  Summary: Exponential families for JAX
5
5
  Author: Neil Girdhar
6
6
  Author-email: Neil Girdhar <mistersheik@gmail.com>
@@ -23,7 +23,7 @@ Requires-Dist: numpy>=2.0
23
23
  Requires-Dist: opt-einsum>=3.4
24
24
  Requires-Dist: optimistix>=0.0.9
25
25
  Requires-Dist: optype[numpy]>=0.8.0
26
- Requires-Dist: tjax>=1.4.4
26
+ Requires-Dist: tjax>=1.4.6
27
27
  Requires-Dist: typing-extensions>=4.8
28
28
  Maintainer: Neil Girdhar
29
29
  Maintainer-email: Neil Girdhar <mistersheik@gmail.com>
@@ -50,6 +50,12 @@ Description-Content-Type: text/x-rst
50
50
  :target: https://scientific-python.org/specs/spec-0000/
51
51
  :alt: SPEC-0
52
52
  :align: center
53
+ .. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json
54
+ :alt: Ruff
55
+ :target: https://github.com/astral-sh/ruff
56
+ .. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ty/main/assets/badge/v0.json
57
+ :alt: ty
58
+ :target: https://github.com/astral-sh/ty
53
59
  .. image:: https://img.shields.io/pypi/pyversions/efax
54
60
  :alt: PyPI - Python Version
55
61
  :align: center
@@ -161,7 +167,10 @@ Every `NaturalParametrization` has methods:
161
167
  - `jeffreys_prior_density`, which returns the square root of the Fisher information
162
168
  determinant,
163
169
  - `characteristic_function`, which evaluates the characteristic function of the sufficient
164
- statistics via analytic continuation of the log-normalizer, and
170
+ statistics via analytic continuation of the log-normalizer (note: for distributions such
171
+ as `GammaNP` that keep one or more natural parameters real in `_complexify`, the result
172
+ is approximate — check `NaturalParametrization.characteristic_function_exact` on the
173
+ class before relying on every component), and
165
174
  - `kl_divergence`, which is the KL divergence.
166
175
 
167
176
  Every `ExpectationParametrization` has methods:
@@ -169,6 +178,11 @@ Every `ExpectationParametrization` has methods:
169
178
  - `to_nat` to convert itself to natural parameters, and
170
179
  - `kl_divergence`, which is the KL divergence.
171
180
 
181
+ The module-level function `expectation_parameters_from_characteristic_function(t, cf_values)`
182
+ estimates the expectation parameters from characteristic-function evaluations via ordinary
183
+ least squares. It is exact for the normal distribution (whose log-CF is linear in the
184
+ frequencies) and a first-order approximation for other families.
185
+
172
186
  Some parametrizations inherit from these interfaces:
173
187
 
174
188
  - `HasConjugatePrior` can produce and recover the conjugate prior,
@@ -313,7 +327,7 @@ EFAX supports the following distributions:
313
327
  - Dirichlet
314
328
  - generalized Dirichlet
315
329
 
316
- - on the n-sphere:
330
+ - on spheres:
317
331
 
318
332
  - von Mises-Fisher
319
333
  - complex von Mises
@@ -448,17 +462,23 @@ Using the cross entropy to iteratively optimize a prediction is simple:
448
462
  return x - 1e-4 * x_bar
449
463
 
450
464
 
451
- def body_fun(q: BernoulliNP) -> BernoulliNP:
452
- q_bar = gradient_cross_entropy(target_distribution, q)
453
- return parameter_map(apply, q, q_bar)
465
+ # Carry (q, q_bar) as the loop state so the gradient is computed exactly
466
+ # once per step rather than once in cond_fun and once in body_fun.
467
+ type State = tuple[BernoulliNP, BernoulliNP]
454
468
 
455
469
 
456
- def cond_fun(q: BernoulliNP) -> JaxBooleanArray:
457
- q_bar = gradient_cross_entropy(target_distribution, q)
470
+ def cond_fun(state: State) -> JaxBooleanArray:
471
+ _, q_bar = state
458
472
  total = jnp.sum(parameter_dot_product(q_bar, q_bar))
459
473
  return total > 1e-6 # noqa: PLR2004
460
474
 
461
475
 
476
+ def body_fun(state: State) -> State:
477
+ q, q_bar = state
478
+ q_new = parameter_map(apply, q, q_bar)
479
+ return q_new, gradient_cross_entropy(target_distribution, q_new)
480
+
481
+
462
482
  # The target_distribution is represented as the expectation parameters of a
463
483
  # Bernoulli distribution corresponding to probabilities 0.3, 0.4, and 0.7.
464
484
  target_distribution = BernoulliEP(jnp.asarray([0.3, 0.4, 0.7]))
@@ -467,10 +487,13 @@ Using the cross entropy to iteratively optimize a prediction is simple:
467
487
  # of a Bernoulli distribution corresponding to log-odds 0, which is probability
468
488
  # 0.5.
469
489
  initial_predictive_distribution = BernoulliNP(jnp.zeros(3))
490
+ initial_gradient = gradient_cross_entropy(target_distribution,
491
+ initial_predictive_distribution)
470
492
 
471
493
  # Optimize the predictive distribution iteratively.
472
- predictive_distribution = lax.while_loop(cond_fun, body_fun,
473
- initial_predictive_distribution)
494
+ predictive_distribution, _ = lax.while_loop(
495
+ cond_fun, body_fun, (initial_predictive_distribution, initial_gradient)
496
+ )
474
497
 
475
498
  # Compare the optimized predictive distribution with the target value in the
476
499
  # same natural parametrization.
@@ -535,6 +558,8 @@ instead.
535
558
  # Dirichlet distribution.
536
559
 
537
560
  # Take the mean over the first axis.
561
+ # parameter_mean averages only variable parameters; fixed parameters (e.g.
562
+ # the failure count in NegativeBinomialNP) are preserved unchanged.
538
563
  ss_mean = parameter_mean(ss, axis=0) # ss_mean also has type DirichletEP.
539
564
 
540
565
  # Convert this back to the natural parametrization.
@@ -17,6 +17,12 @@
17
17
  :target: https://scientific-python.org/specs/spec-0000/
18
18
  :alt: SPEC-0
19
19
  :align: center
20
+ .. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json
21
+ :alt: Ruff
22
+ :target: https://github.com/astral-sh/ruff
23
+ .. image:: https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ty/main/assets/badge/v0.json
24
+ :alt: ty
25
+ :target: https://github.com/astral-sh/ty
20
26
  .. image:: https://img.shields.io/pypi/pyversions/efax
21
27
  :alt: PyPI - Python Version
22
28
  :align: center
@@ -128,7 +134,10 @@ Every `NaturalParametrization` has methods:
128
134
  - `jeffreys_prior_density`, which returns the square root of the Fisher information
129
135
  determinant,
130
136
  - `characteristic_function`, which evaluates the characteristic function of the sufficient
131
- statistics via analytic continuation of the log-normalizer, and
137
+ statistics via analytic continuation of the log-normalizer (note: for distributions such
138
+ as `GammaNP` that keep one or more natural parameters real in `_complexify`, the result
139
+ is approximate — check `NaturalParametrization.characteristic_function_exact` on the
140
+ class before relying on every component), and
132
141
  - `kl_divergence`, which is the KL divergence.
133
142
 
134
143
  Every `ExpectationParametrization` has methods:
@@ -136,6 +145,11 @@ Every `ExpectationParametrization` has methods:
136
145
  - `to_nat` to convert itself to natural parameters, and
137
146
  - `kl_divergence`, which is the KL divergence.
138
147
 
148
+ The module-level function `expectation_parameters_from_characteristic_function(t, cf_values)`
149
+ estimates the expectation parameters from characteristic-function evaluations via ordinary
150
+ least squares. It is exact for the normal distribution (whose log-CF is linear in the
151
+ frequencies) and a first-order approximation for other families.
152
+
139
153
  Some parametrizations inherit from these interfaces:
140
154
 
141
155
  - `HasConjugatePrior` can produce and recover the conjugate prior,
@@ -280,7 +294,7 @@ EFAX supports the following distributions:
280
294
  - Dirichlet
281
295
  - generalized Dirichlet
282
296
 
283
- - on the n-sphere:
297
+ - on spheres:
284
298
 
285
299
  - von Mises-Fisher
286
300
  - complex von Mises
@@ -415,17 +429,23 @@ Using the cross entropy to iteratively optimize a prediction is simple:
415
429
  return x - 1e-4 * x_bar
416
430
 
417
431
 
418
- def body_fun(q: BernoulliNP) -> BernoulliNP:
419
- q_bar = gradient_cross_entropy(target_distribution, q)
420
- return parameter_map(apply, q, q_bar)
432
+ # Carry (q, q_bar) as the loop state so the gradient is computed exactly
433
+ # once per step rather than once in cond_fun and once in body_fun.
434
+ type State = tuple[BernoulliNP, BernoulliNP]
421
435
 
422
436
 
423
- def cond_fun(q: BernoulliNP) -> JaxBooleanArray:
424
- q_bar = gradient_cross_entropy(target_distribution, q)
437
+ def cond_fun(state: State) -> JaxBooleanArray:
438
+ _, q_bar = state
425
439
  total = jnp.sum(parameter_dot_product(q_bar, q_bar))
426
440
  return total > 1e-6 # noqa: PLR2004
427
441
 
428
442
 
443
+ def body_fun(state: State) -> State:
444
+ q, q_bar = state
445
+ q_new = parameter_map(apply, q, q_bar)
446
+ return q_new, gradient_cross_entropy(target_distribution, q_new)
447
+
448
+
429
449
  # The target_distribution is represented as the expectation parameters of a
430
450
  # Bernoulli distribution corresponding to probabilities 0.3, 0.4, and 0.7.
431
451
  target_distribution = BernoulliEP(jnp.asarray([0.3, 0.4, 0.7]))
@@ -434,10 +454,13 @@ Using the cross entropy to iteratively optimize a prediction is simple:
434
454
  # of a Bernoulli distribution corresponding to log-odds 0, which is probability
435
455
  # 0.5.
436
456
  initial_predictive_distribution = BernoulliNP(jnp.zeros(3))
457
+ initial_gradient = gradient_cross_entropy(target_distribution,
458
+ initial_predictive_distribution)
437
459
 
438
460
  # Optimize the predictive distribution iteratively.
439
- predictive_distribution = lax.while_loop(cond_fun, body_fun,
440
- initial_predictive_distribution)
461
+ predictive_distribution, _ = lax.while_loop(
462
+ cond_fun, body_fun, (initial_predictive_distribution, initial_gradient)
463
+ )
441
464
 
442
465
  # Compare the optimized predictive distribution with the target value in the
443
466
  # same natural parametrization.
@@ -502,6 +525,8 @@ instead.
502
525
  # Dirichlet distribution.
503
526
 
504
527
  # Take the mean over the first axis.
528
+ # parameter_mean averages only variable parameters; fixed parameters (e.g.
529
+ # the failure count in NegativeBinomialNP) are preserved unchanged.
505
530
  ss_mean = parameter_mean(ss, axis=0) # ss_mean also has type DirichletEP.
506
531
 
507
532
  # Convert this back to the natural parametrization.
@@ -108,7 +108,8 @@ class CategoricalEP(
108
108
  """The expectation parametrization of the categorical distribution.
109
109
 
110
110
  Args:
111
- probability: The probability vector with the final element omitted, i.e., [p_i]_{i in 1...n-1}.
111
+ probability: The probability vector with the final element omitted, i.e.,
112
+ [p_i]_{i in 1...n-1}.
112
113
  """
113
114
 
114
115
  probability: JaxRealArray = distribution_parameter(
@@ -1,6 +1,6 @@
1
1
  from __future__ import annotations
2
2
 
3
- from typing import Self, override
3
+ from typing import ClassVar, Self, override
4
4
 
5
5
  import jax.random as jr
6
6
  import jax.scipy.special as jss
@@ -39,6 +39,8 @@ class GammaNP(
39
39
  shape_minus_one: The shape minus one.
40
40
  """
41
41
 
42
+ characteristic_function_exact: ClassVar[bool] = False
43
+
42
44
  negative_rate: JaxRealArray = distribution_parameter(ScalarSupport(ring=negative_support))
43
45
  shape_minus_one: JaxRealArray = distribution_parameter(
44
46
  ScalarSupport(ring=RealField(minimum=-1.0))
@@ -83,9 +85,10 @@ class GammaNP(
83
85
  def _complexify(self, t: Self) -> Self: # ty: ignore
84
86
  """Only complexify negative_rate; jss.gammaln does not accept complex inputs.
85
87
 
86
- This means characteristic_function cannot query the log x sufficient
87
- statistic (the shape_minus_one component). Querying it returns 1
88
- instead of the true E[exp(i·t·log x)].
88
+ shape_minus_one is kept real, so ``characteristic_function`` cannot
89
+ query the log-x sufficient statistic: that component returns 1 instead
90
+ of the true E[exp(i·t·log x)]. Consequently,
91
+ ``characteristic_function_exact = False`` on this class.
89
92
  """
90
93
  return type(self)(self.negative_rate + 1j * t.negative_rate, self.shape_minus_one)
91
94
 
@@ -55,7 +55,8 @@ 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), max_steps=1000
58
+ solver=optx.Newton[JaxRealArray, JaxRealArray, None](rtol=0.0, atol=1e-7), # ty: ignore
59
+ max_steps=1000,
59
60
  )
60
61
 
61
62
 
@@ -1,7 +1,7 @@
1
1
  from __future__ import annotations
2
2
 
3
3
  from abc import abstractmethod
4
- from typing import TYPE_CHECKING, Any, Generic, Self, final, get_type_hints
4
+ from typing import TYPE_CHECKING, Any, ClassVar, Generic, Self, final, get_type_hints
5
5
 
6
6
  import jax
7
7
  import jax.numpy as jnp
@@ -53,8 +53,18 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
53
53
 
54
54
  The motivation for the natural parametrization is combining and scaling independent predictive
55
55
  evidence. In the natural parametrization, these operations correspond to scaling and addition.
56
+
57
+ Class variables:
58
+ characteristic_function_exact: True when `_complexify` shifts every natural parameter
59
+ into the complex plane, so `characteristic_function` queries all sufficient
60
+ statistics exactly. Set to False in subclasses that must keep one or more
61
+ parameters real (e.g. because `log_normalizer` calls a function such as
62
+ ``gammaln`` that does not accept complex inputs); those components return 1
63
+ instead of the true characteristic-function value.
56
64
  """
57
65
 
66
+ characteristic_function_exact: ClassVar[bool] = True
67
+
58
68
  @abstract_custom_jvp(_log_normalizer_jvp)
59
69
  @abstract_jit
60
70
  @abstractmethod
@@ -196,11 +206,13 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
196
206
  def _complexify(self, t: Self) -> Self:
197
207
  """Shift natural parameters into the complex plane: η → η + i·t.
198
208
 
199
- The default shifts every field. Subclasses must override this when
200
- ``log_normalizer`` calls functions that don't accept complex inputs
201
- (e.g. ``gammaln``). Keeping parameter k real means
202
- ``characteristic_function`` cannot query the k-th sufficient statistic:
203
- it will silently return 1 for that component instead of the true value.
209
+ The default shifts every field, which is correct for any distribution
210
+ whose ``log_normalizer`` accepts fully complex inputs. Subclasses that
211
+ must keep one or more parameters real (e.g. because ``log_normalizer``
212
+ calls ``gammaln``, which rejects complex inputs) must override this
213
+ method **and** set ``characteristic_function_exact = False`` on the
214
+ class. For those parameters, ``characteristic_function`` returns 1
215
+ instead of the true E[exp(i·⟨t, T(x)⟩)] component.
204
216
 
205
217
  Args:
206
218
  t: Imaginary displacements, same pytree structure as self.
@@ -237,6 +249,13 @@ class NaturalParametrization(Distribution, JaxAbstractClass, Generic[EP, Domain]
237
249
  t_grid = PoissonNP(omegas) # shape (k,)
238
250
  phi = p.characteristic_function(t_grid) # shape (k,)
239
251
 
252
+ Note: when ``type(self).characteristic_function_exact`` is False, one
253
+ or more natural parameters are kept real in ``_complexify`` because
254
+ ``log_normalizer`` uses a function that rejects complex inputs. The
255
+ corresponding sufficient-statistic components return 1 rather than
256
+ the true value. Check the class attribute before relying on the
257
+ result for those components.
258
+
240
259
  Args:
241
260
  t: Frequencies in natural-parameter space, same structure as self.
242
261
 
@@ -31,7 +31,6 @@ def triangular_number_index(k: int, /) -> int | None:
31
31
 
32
32
 
33
33
  class Support:
34
- @override
35
34
  def __init__(self, *, ring: Ring = real_field) -> None:
36
35
  super().__init__()
37
36
  self.ring = ring
@@ -280,6 +279,12 @@ class SquareMatrixSupport(Support):
280
279
  x = self.ring.unflattened(y, map_from_plane=map_from_plane)
281
280
  return xp.reshape(x, x.shape[:-1] + self.shape(dimensions))
282
281
 
282
+ @override
283
+ def generate(
284
+ self, xp: Namespace, rng: Generator, shape: Shape, safety: float, dimensions: int
285
+ ) -> JaxRealArray:
286
+ return self.ring.generate(xp, rng, (*shape, dimensions, dimensions), safety)
287
+
283
288
 
284
289
  class CircularBoundedSupport(VectorSupport):
285
290
  def __init__(self, radius: float) -> None:
@@ -68,7 +68,9 @@ class Assembler(Generic[P]):
68
68
 
69
69
  infos = []
70
70
  for info in self.infos:
71
- assert issubclass(info.type_, ExpectationParametrization)
71
+ if not issubclass(info.type_, ExpectationParametrization):
72
+ msg = f"{info.type_.__name__} is not an EP"
73
+ raise TypeError(msg)
72
74
  infos.append(
73
75
  SubDistributionInfo(
74
76
  info.path,
@@ -85,7 +87,9 @@ class Assembler(Generic[P]):
85
87
 
86
88
  infos: list[SubDistributionInfo] = []
87
89
  for info in self.infos:
88
- assert issubclass(info.type_, NaturalParametrization)
90
+ if not issubclass(info.type_, NaturalParametrization):
91
+ msg = f"{info.type_.__name__} is not an NP"
92
+ raise TypeError(msg)
89
93
  infos.append(
90
94
  SubDistributionInfo(
91
95
  info.path,
@@ -129,7 +133,7 @@ class Assembler(Generic[P]):
129
133
  ]
130
134
  q_values = parameters(q).values()
131
135
  q_params_as_p = dict(zip(p_paths, q_values, strict=True))
132
- return self.assemble(q_params_as_p)
136
+ return self.assemble(q_params_as_p) # ty: ignore
133
137
 
134
138
  def domain_support(self) -> dict[Path, Support]:
135
139
  """Return the domain support constraints for each simple sub-distribution in the tree."""
@@ -45,7 +45,9 @@ class Estimator(Assembler[P]):
45
45
  ExpectationParametrization,
46
46
  )
47
47
 
48
- assert issubclass(type_p, ExpectationParametrization)
48
+ if not issubclass(type_p, ExpectationParametrization):
49
+ msg = f"{type_p.__name__} is not an EP"
50
+ raise TypeError(msg)
49
51
  return Estimator(
50
52
  [SubDistributionInfo((), type_p, 0, [])],
51
53
  {(name,): value for name, value in fixed_parameters.items()},
@@ -63,7 +65,9 @@ class Estimator(Assembler[P]):
63
65
  )
64
66
 
65
67
  infos = cls.create_assembler(p).infos
66
- assert isinstance(p, ExpectationParametrization)
68
+ if not isinstance(p, ExpectationParametrization):
69
+ msg = f"{type(p).__name__} is not an EP"
70
+ raise TypeError(msg)
67
71
  fixed_parameters = parameters(p, fixed=True)
68
72
  return cls(infos, fixed_parameters)
69
73
 
@@ -4,7 +4,7 @@ from collections import defaultdict
4
4
  from collections.abc import Callable, Iterable, Mapping
5
5
  from functools import reduce
6
6
  from itertools import starmap
7
- from typing import TYPE_CHECKING, TypeVar
7
+ from typing import TypeVar
8
8
 
9
9
  from array_api_compat import array_namespace
10
10
  from jax import jit
@@ -17,7 +17,7 @@ from .types import Axis
17
17
 
18
18
 
19
19
  @jit
20
- def parameter_dot_product(x: NaturalParametrization, y: Distribution, /) -> JaxRealArray:
20
+ def parameter_dot_product(x: Distribution, y: Distribution, /) -> JaxRealArray:
21
21
  """Return the vectorized dot product over all of the variable parameters."""
22
22
 
23
23
  def dotted_fields() -> Iterable[JaxRealArray]:
@@ -38,12 +38,13 @@ T = TypeVar("T", bound=Distribution)
38
38
 
39
39
 
40
40
  def parameter_mean[T: Distribution](x: T, /, *, axis: Axis | None = None) -> T:
41
- """Return the mean of the parameters (fixed and variable)."""
41
+ """Return the mean of the variable parameters; fixed parameters are preserved unchanged."""
42
42
  xp = array_namespace(x)
43
43
  structure = Assembler.create_assembler(x)
44
- p = parameters(x)
45
- q = {path: xp.mean(value, axis=axis) for path, value in p.items()}
46
- return structure.assemble(q)
44
+ fixed = parameters(x, fixed=True)
45
+ variable = parameters(x, fixed=False)
46
+ averaged = {path: xp.mean(value, axis=axis) for path, value in variable.items()}
47
+ return structure.assemble({**fixed, **averaged})
47
48
 
48
49
 
49
50
  def parameter_map[T: Distribution](
@@ -80,7 +81,3 @@ def _parameter_dot_product(x: JaxComplexArray, y: JaxComplexArray, n_axes: int)
80
81
  xp = array_namespace(x, y)
81
82
  axes = tuple(range(-n_axes, 0))
82
83
  return xp.sum(xp.real(x * xp.conj(y)), axis=axes)
83
-
84
-
85
- if TYPE_CHECKING:
86
- from .natural_parametrization import NaturalParametrization
@@ -50,9 +50,19 @@ class JointDistribution(Distribution):
50
50
  @property
51
51
  @override
52
52
  def shape(self) -> Shape:
53
- for distribution in self._sub_distributions.values():
54
- return distribution.shape
55
- raise ValueError
53
+ shapes = [d.shape for d in self._sub_distributions.values()]
54
+ if not shapes:
55
+ msg = "JointDistribution has no sub-distributions"
56
+ raise ValueError(msg)
57
+ shape = shapes[0]
58
+ if any(s != shape for s in shapes[1:]):
59
+ names = list(self._sub_distributions)
60
+ msg = (
61
+ f"Sub-distribution shapes are inconsistent: "
62
+ f"{dict(zip(names, shapes, strict=True))}"
63
+ )
64
+ raise ValueError(msg)
65
+ return shape
56
66
 
57
67
 
58
68
  @dataclass
@@ -6,19 +6,18 @@ build-backend = "uv_build"
6
6
  dev = [
7
7
  "jupyter>=1",
8
8
  "lefthook>=2",
9
- "pyright>=1.1.408",
10
9
  "pytest-xdist[psutil]>=3",
11
10
  "pytest>=9",
12
11
  "ruff>=0.15",
13
12
  "scipy-stubs>=1.15",
14
13
  "scipy>=1.15",
15
14
  "toml-sort>=0.24",
16
- "ty>=0.0.22"
15
+ "ty>=0.0.29"
17
16
  ]
18
17
 
19
18
  [project]
20
19
  name = "efax"
21
- version = "2.2.0"
20
+ version = "2.2.4"
22
21
  description = "Exponential families for JAX"
23
22
  readme = "README.rst"
24
23
  requires-python = ">=3.12, <3.15"
@@ -46,7 +45,7 @@ dependencies = [
46
45
  "opt-einsum>=3.4",
47
46
  "optimistix>=0.0.9",
48
47
  "optype[numpy]>=0.8.0",
49
- "tjax>=1.4.4",
48
+ "tjax>=1.4.6",
50
49
  "typing_extensions>=4.8"
51
50
  ]
52
51
 
@@ -2,6 +2,7 @@
2
2
 
3
3
  from ._src.distributions.bernoulli import BernoulliEP, BernoulliNP
4
4
  from ._src.distributions.beta import BetaEP, BetaNP
5
+ from ._src.distributions.categorical import CategoricalEP, CategoricalNP
5
6
  from ._src.distributions.chi import ChiEP, ChiNP
6
7
  from ._src.distributions.chi_square import ChiSquareEP, ChiSquareNP
7
8
  from ._src.distributions.cmvn.circularly_symmetric import (
@@ -31,7 +32,6 @@ from ._src.distributions.log_normal.unit_variance import (
31
32
  UnitVarianceLogNormalNP,
32
33
  )
33
34
  from ._src.distributions.logarithmic import LogarithmicEP, LogarithmicNP
34
- from ._src.distributions.categorical import CategoricalEP, CategoricalNP
35
35
  from ._src.distributions.multivariate_normal.arbitrary import (
36
36
  MultivariateNormalEP,
37
37
  MultivariateNormalNP,
@@ -105,6 +105,8 @@ __all__ = [
105
105
  "BetaEP",
106
106
  "BetaNP",
107
107
  "BooleanRing",
108
+ "CategoricalEP",
109
+ "CategoricalNP",
108
110
  "ChiEP",
109
111
  "ChiNP",
110
112
  "ChiSquareEP",
@@ -156,8 +158,6 @@ __all__ = [
156
158
  "LogarithmicEP",
157
159
  "LogarithmicNP",
158
160
  "Multidimensional",
159
- "CategoricalEP",
160
- "CategoricalNP",
161
161
  "MultivariateDiagonalNormalEP",
162
162
  "MultivariateDiagonalNormalNP",
163
163
  "MultivariateDiagonalNormalVP",
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes