efax 2.2.4__tar.gz → 2.2.5__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.2.5}/PKG-INFO +1 -1
  2. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/fixed_variance.py +4 -1
  3. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/weibull.py +2 -2
  4. {efax-2.2.4 → efax-2.2.5}/efax/_src/tools.py +9 -5
  5. {efax-2.2.4 → efax-2.2.5}/pyproject.toml +1 -1
  6. {efax-2.2.4 → efax-2.2.5}/README.rst +0 -0
  7. {efax-2.2.4 → efax-2.2.5}/efax/__init__.py +0 -0
  8. {efax-2.2.4 → efax-2.2.5}/efax/_src/__init__.py +0 -0
  9. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/__init__.py +0 -0
  10. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/bernoulli.py +0 -0
  11. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/beta.py +0 -0
  12. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/categorical.py +0 -0
  13. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/chi.py +0 -0
  14. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/chi_square.py +0 -0
  15. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/__init__.py +0 -0
  16. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
  17. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
  18. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/__init__.py +0 -0
  19. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
  20. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
  21. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_von_mises.py +0 -0
  22. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/dirichlet.py +0 -0
  23. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/dirichlet_common.py +0 -0
  24. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/exponential.py +0 -0
  25. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/gamma.py +0 -0
  26. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/gen_dirichlet.py +0 -0
  27. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/geometric.py +0 -0
  28. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/inverse_gamma.py +0 -0
  29. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/inverse_gaussian.py +0 -0
  30. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/__init__.py +0 -0
  31. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/log_normal.py +0 -0
  32. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/unit_variance.py +0 -0
  33. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/logarithmic.py +0 -0
  34. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
  35. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
  36. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
  37. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
  38. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
  39. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/negative_binomial.py +0 -0
  40. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/negative_binomial_common.py +0 -0
  41. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/__init__.py +0 -0
  42. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/normal.py +0 -0
  43. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/unit_variance.py +0 -0
  44. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/poisson.py +0 -0
  45. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/rayleigh.py +0 -0
  46. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
  47. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
  48. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
  49. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/von_mises.py +0 -0
  50. {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/wishart.py +0 -0
  51. {efax-2.2.4 → efax-2.2.5}/efax/_src/expectation_parametrization.py +0 -0
  52. {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/__init__.py +0 -0
  53. {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/conjugate_prior.py +0 -0
  54. {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/multidimensional.py +0 -0
  55. {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/samplable.py +0 -0
  56. {efax-2.2.4 → efax-2.2.5}/efax/_src/iteration.py +0 -0
  57. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/__init__.py +0 -0
  58. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
  59. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +0 -0
  60. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/optimistix.py +0 -0
  61. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/has_entropy.py +0 -0
  62. {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/transformed_parametrization.py +0 -0
  63. {efax-2.2.4 → efax-2.2.5}/efax/_src/natural_parametrization.py +0 -0
  64. {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/__init__.py +0 -0
  65. {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/parameter.py +0 -0
  66. {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/ring.py +0 -0
  67. {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/support.py +0 -0
  68. {efax-2.2.4 → efax-2.2.5}/efax/_src/parametrization.py +0 -0
  69. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/__init__.py +0 -0
  70. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/assembler.py +0 -0
  71. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/estimator.py +0 -0
  72. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/flattener.py +0 -0
  73. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/parameter_names.py +0 -0
  74. {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/parameter_supports.py +0 -0
  75. {efax-2.2.4 → efax-2.2.5}/efax/_src/transform/__init__.py +0 -0
  76. {efax-2.2.4 → efax-2.2.5}/efax/_src/transform/joint.py +0 -0
  77. {efax-2.2.4 → efax-2.2.5}/efax/_src/types.py +0 -0
  78. {efax-2.2.4 → efax-2.2.5}/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.2.5
4
4
  Summary: Exponential families for JAX
5
5
  Author: Neil Girdhar
6
6
  Author-email: Neil Girdhar <mistersheik@gmail.com>
@@ -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:
@@ -38,13 +38,17 @@ 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 variable parameters; fixed parameters are preserved unchanged."""
41
+ """Return the mean over axis of all parameters, including fixed ones.
42
+
43
+ Fixed parameters are constant along the reduced axes (having been broadcast to the sample shape
44
+ in sufficient_statistics), so their mean equals their value. Integer-typed parameters remain
45
+ integers after the mean because the Array API spec requires mean to preserve integer dtypes.
46
+ """
42
47
  xp = array_namespace(x)
43
48
  structure = Assembler.create_assembler(x)
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})
49
+ all_params = parameters(x)
50
+ averaged = {path: xp.mean(value, axis=axis) for path, value in all_params.items()}
51
+ return structure.assemble(averaged)
48
52
 
49
53
 
50
54
  def parameter_map[T: Distribution](
@@ -17,7 +17,7 @@ dev = [
17
17
 
18
18
  [project]
19
19
  name = "efax"
20
- version = "2.2.4"
20
+ version = "2.2.5"
21
21
  description = "Exponential families for JAX"
22
22
  readme = "README.rst"
23
23
  requires-python = ">=3.12, <3.15"
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes