Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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)
85 changes: 85 additions & 0 deletions tests/distributions/test_custom_factory_subclass_preservation.py
Original file line number Diff line number Diff line change
@@ -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()
Loading