diff --git a/src/pyrecest/distributions/cart_prod/custom_hypercylindrical_distribution.py b/src/pyrecest/distributions/cart_prod/custom_hypercylindrical_distribution.py index ac75c0860..8df5ed0d2 100644 --- a/src/pyrecest/distributions/cart_prod/custom_hypercylindrical_distribution.py +++ b/src/pyrecest/distributions/cart_prod/custom_hypercylindrical_distribution.py @@ -28,8 +28,8 @@ def __init__(self, f, bound_dim, lin_dim): self, f, bound_dim, lin_dim ) - @staticmethod - def from_distribution(distribution): + @classmethod + def from_distribution(cls, distribution): """ Create a CustomHypercylindricalDistribution from another AbstractHypercylindricalDistribution. @@ -41,10 +41,7 @@ def from_distribution(distribution): chhd (CustomHypercylindricalDistribution) The created CustomHypercylindricalDistribution """ - chhd = CustomHypercylindricalDistribution( - distribution.pdf, distribution.bound_dim, distribution.lin_dim - ) - return chhd + return cls(distribution.pdf, distribution.bound_dim, distribution.lin_dim) def integrate(self, integration_boundaries=None): # Call the integrate method from the superclass diff --git a/src/pyrecest/distributions/hypersphere_subset/custom_hemispherical_distribution.py b/src/pyrecest/distributions/hypersphere_subset/custom_hemispherical_distribution.py index fb691191e..1c950edad 100644 --- a/src/pyrecest/distributions/hypersphere_subset/custom_hemispherical_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/custom_hemispherical_distribution.py @@ -17,15 +17,15 @@ def __init__(self, f: Callable): AbstractHemisphericalDistribution.__init__(self) CustomHyperhemisphericalDistribution.__init__(self, f, 2) - @staticmethod - def from_distribution(distribution: "AbstractHypersphericalDistribution"): + @classmethod + def from_distribution(cls, distribution: "AbstractHypersphericalDistribution"): if distribution.dim != 2: raise ValueError("Dimension of the distribution should be 2.") if isinstance(distribution, AbstractHyperhemisphericalDistribution): - return CustomHemisphericalDistribution(distribution.pdf) + return cls(distribution.pdf) if isinstance(distribution, BinghamDistribution): - chsd = CustomHemisphericalDistribution(distribution.pdf) + chsd = cls(distribution.pdf) chsd.scale_by = 2 return chsd if isinstance(distribution, AbstractHypersphericalDistribution): @@ -38,7 +38,7 @@ def from_distribution(distribution: "AbstractHypersphericalDistribution"): distribution.pdf, distribution.dim ) norm_const_inv = chhd_unnorm.integrate() - chsd = CustomHemisphericalDistribution(distribution.pdf) + chsd = cls(distribution.pdf) chsd.scale_by = 1 / norm_const_inv return chsd diff --git a/src/pyrecest/distributions/hypersphere_subset/custom_hyperhemispherical_distribution.py b/src/pyrecest/distributions/hypersphere_subset/custom_hyperhemispherical_distribution.py index 3e1fdcf42..b73d1d972 100644 --- a/src/pyrecest/distributions/hypersphere_subset/custom_hyperhemispherical_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/custom_hyperhemispherical_distribution.py @@ -52,8 +52,8 @@ def integrate(self, integration_boundaries=None): self, integration_boundaries ) - @staticmethod - def from_distribution(distribution: "AbstractHypersphericalDistribution"): + @classmethod + def from_distribution(cls, distribution: "AbstractHypersphericalDistribution"): """ Create a CustomHyperhemisphericalDistribution from another distribution. @@ -62,21 +62,15 @@ def from_distribution(distribution: "AbstractHypersphericalDistribution"): :raises ValueError: if the type of dist is not supported. """ if isinstance(distribution, AbstractHyperhemisphericalDistribution): - return CustomHyperhemisphericalDistribution( - distribution.pdf, distribution.dim - ) + return cls(distribution.pdf, distribution.dim) if isinstance(distribution, BinghamDistribution): - chhd = CustomHyperhemisphericalDistribution( - distribution.pdf, distribution.dim - ) + chhd = cls(distribution.pdf, distribution.dim) chhd.scale_by = 2 return chhd if isinstance(distribution, AbstractHypersphericalDistribution): - chhd = CustomHyperhemisphericalDistribution( - distribution.pdf, distribution.dim - ) + chhd = cls(distribution.pdf, distribution.dim) norm_const_inv = chhd.integrate() chhd.scale_by = 1 / norm_const_inv return chhd diff --git a/src/pyrecest/distributions/hypersphere_subset/custom_hyperspherical_distribution.py b/src/pyrecest/distributions/hypersphere_subset/custom_hyperspherical_distribution.py index 7bb8c2660..5fc8162ac 100644 --- a/src/pyrecest/distributions/hypersphere_subset/custom_hyperspherical_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/custom_hyperspherical_distribution.py @@ -9,13 +9,12 @@ def __init__(self, f, dim, scale_by=1): AbstractCustomDistribution.__init__(self, f, scale_by) AbstractHypersphericalDistribution.__init__(self, dim) - @staticmethod - def from_distribution(distribution): + @classmethod + def from_distribution(cls, distribution): if not isinstance(distribution, AbstractHypersphericalDistribution): raise ValueError("Input variable distribution is of the wrong class.") - chd = CustomHypersphericalDistribution(distribution.pdf, distribution.dim) - return chd + return cls(distribution.pdf, distribution.dim) def integrate(self, integration_boundaries=None): return AbstractHypersphericalDistribution.integrate( diff --git a/src/pyrecest/distributions/nonperiodic/custom_linear_distribution.py b/src/pyrecest/distributions/nonperiodic/custom_linear_distribution.py index aac90a891..be8e888f3 100644 --- a/src/pyrecest/distributions/nonperiodic/custom_linear_distribution.py +++ b/src/pyrecest/distributions/nonperiodic/custom_linear_distribution.py @@ -83,8 +83,8 @@ def pdf(self, xs): p = reshape(p, xs.shape[:-1]) return p - @staticmethod - def from_distribution(distribution): + @classmethod + def from_distribution(cls, distribution): """ Creates a CustomLinearDistribution from some other distribution @@ -96,8 +96,7 @@ def from_distribution(distribution): chd (CustomLinearDistribution) CustomLinearDistribution with identical pdf """ - chd = CustomLinearDistribution(distribution.pdf, distribution.dim) - return chd + return cls(distribution.pdf, distribution.dim) def integrate(self, left=None, right=None): return AbstractLinearDistribution.integrate(self, left, right) diff --git a/tests/distributions/test_custom_factory_subclass_preservation.py b/tests/distributions/test_custom_factory_subclass_preservation.py new file mode 100644 index 000000000..f42981a9d --- /dev/null +++ b/tests/distributions/test_custom_factory_subclass_preservation.py @@ -0,0 +1,85 @@ +import unittest + +from pyrecest.distributions.cart_prod.custom_hypercylindrical_distribution import ( + CustomHypercylindricalDistribution, +) +from pyrecest.distributions.conversion import convert_distribution +from pyrecest.distributions.hypersphere_subset.custom_hemispherical_distribution import ( + CustomHemisphericalDistribution, +) +from pyrecest.distributions.hypersphere_subset.custom_hyperhemispherical_distribution import ( + CustomHyperhemisphericalDistribution, +) +from pyrecest.distributions.hypersphere_subset.custom_hyperspherical_distribution import ( + CustomHypersphericalDistribution, +) +from pyrecest.distributions.nonperiodic.custom_linear_distribution import ( + CustomLinearDistribution, +) + + +def _constant_pdf(xs): + return xs[..., 0] * 0.0 + 1.0 + + +class _CustomLinearSubclass(CustomLinearDistribution): + pass + + +class _CustomHypercylindricalSubclass(CustomHypercylindricalDistribution): + pass + + +class _CustomHypersphericalSubclass(CustomHypersphericalDistribution): + pass + + +class _CustomHyperhemisphericalSubclass(CustomHyperhemisphericalDistribution): + pass + + +class _CustomHemisphericalSubclass(CustomHemisphericalDistribution): + pass + + +class CustomFactorySubclassPreservationTest(unittest.TestCase): + def test_custom_linear_conversion_preserves_requested_subclass(self): + source = CustomLinearDistribution(_constant_pdf, dim=1) + + converted = convert_distribution(source, _CustomLinearSubclass) + + self.assertIsInstance(converted, _CustomLinearSubclass) + + def test_custom_hypercylindrical_conversion_preserves_requested_subclass(self): + source = CustomHypercylindricalDistribution( + _constant_pdf, bound_dim=1, lin_dim=1 + ) + + converted = convert_distribution(source, _CustomHypercylindricalSubclass) + + self.assertIsInstance(converted, _CustomHypercylindricalSubclass) + + def test_custom_hyperspherical_conversion_preserves_requested_subclass(self): + source = CustomHypersphericalDistribution(_constant_pdf, dim=2) + + converted = convert_distribution(source, _CustomHypersphericalSubclass) + + self.assertIsInstance(converted, _CustomHypersphericalSubclass) + + def test_custom_hyperhemispherical_conversion_preserves_requested_subclass(self): + source = CustomHyperhemisphericalDistribution(_constant_pdf, dim=2) + + converted = convert_distribution(source, _CustomHyperhemisphericalSubclass) + + self.assertIsInstance(converted, _CustomHyperhemisphericalSubclass) + + def test_custom_hemispherical_conversion_preserves_requested_subclass(self): + source = CustomHemisphericalDistribution(_constant_pdf) + + converted = convert_distribution(source, _CustomHemisphericalSubclass) + + self.assertIsInstance(converted, _CustomHemisphericalSubclass) + + +if __name__ == "__main__": + unittest.main()