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.
- {efax-2.2.4 → efax-2.2.5}/PKG-INFO +1 -1
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/fixed_variance.py +4 -1
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/weibull.py +2 -2
- {efax-2.2.4 → efax-2.2.5}/efax/_src/tools.py +9 -5
- {efax-2.2.4 → efax-2.2.5}/pyproject.toml +1 -1
- {efax-2.2.4 → efax-2.2.5}/README.rst +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/bernoulli.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/beta.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/categorical.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/chi.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/chi_square.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/circularly_symmetric.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/cmvn/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/complex_normal.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/complex_von_mises.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/dirichlet.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/dirichlet_common.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/exponential.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/gamma.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/gen_dirichlet.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/geometric.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/inverse_gamma.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/inverse_gaussian.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/log_normal.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/log_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/logarithmic.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/arbitrary.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/diagonal.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/isotropic.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/multivariate_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/negative_binomial.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/negative_binomial_common.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/normal.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/poisson.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/rayleigh.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/softplus.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/softplus_normal/unit_variance.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/von_mises.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/distributions/wishart.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/expectation_parametrization.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/conjugate_prior.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/multidimensional.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/interfaces/samplable.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/iteration.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/exp_to_nat.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/exp_to_nat/optimistix.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/has_entropy.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/mixins/transformed_parametrization.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/natural_parametrization.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/parameter.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/ring.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/parameter/support.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/parametrization.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/assembler.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/estimator.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/flattener.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/parameter_names.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/structure/parameter_supports.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/transform/__init__.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/transform/joint.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/_src/types.py +0 -0
- {efax-2.2.4 → efax-2.2.5}/efax/py.typed +0 -0
|
@@ -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:
|
|
@@ -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
|
|
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
|
-
|
|
45
|
-
|
|
46
|
-
|
|
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](
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|