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 @@ -357,8 +357,9 @@ def get_full_c(self):
full_c = concatenate([neg_c, self.c]) # Concatenate arrays to get full spectrum
return full_c

@staticmethod
@classmethod
def from_distribution(
cls,
distribution: AbstractCircularDistribution,
n: Union[int, int32, int64],
transformation: str = "sqrt",
Expand All @@ -379,7 +380,7 @@ def from_distribution(
for k in range(int(n) // 2 + 1)
]
)
fd = CircularFourierDistribution(
fd = cls(
c=coeffs,
n=int(n),
transformation=transformation,
Expand All @@ -398,14 +399,15 @@ def from_distribution(
fvals = sqrt(fvals)
else:
raise NotImplementedError("Transformation not supported.")
fd = CircularFourierDistribution.from_function_values(
fd = cls.from_function_values(
fvals, transformation, store_values_multiplied_by_n
)

return fd

@staticmethod
@classmethod
def from_function_values(
cls,
fvals,
transformation: str = "sqrt",
store_values_multiplied_by_n: bool = True,
Expand All @@ -416,7 +418,7 @@ def from_function_values(
if not store_values_multiplied_by_n:
c = c * (1.0 / n_values)

fd = CircularFourierDistribution(
fd = cls(
c=c,
transformation=transformation,
n=n_values,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
import unittest

# pylint: disable=no-name-in-module,no-member
from pyrecest.backend import array
from pyrecest.distributions.circle.circular_dirac_distribution import (
CircularDiracDistribution,
)
from pyrecest.distributions.circle.circular_fourier_distribution import (
CircularFourierDistribution,
)
from pyrecest.distributions.circle.von_mises_distribution import VonMisesDistribution
from pyrecest.distributions.conversion import convert_distribution


class _CircularFourierSubclass(CircularFourierDistribution):
pass


class CircularFourierFactorySubclassPreservationTest(unittest.TestCase):
def test_conversion_factory_preserves_requested_subclass_for_density_source(self):
source = VonMisesDistribution(0.3, 2.0)

converted = convert_distribution(
source,
_CircularFourierSubclass,
n=9,
transformation="identity",
store_values_multiplied_by_n=False,
)

self.assertIsInstance(converted, _CircularFourierSubclass)

def test_conversion_factory_preserves_requested_subclass_for_dirac_source(self):
source = CircularDiracDistribution(array([0.0, 1.0]))

converted = convert_distribution(
source,
_CircularFourierSubclass,
n=9,
transformation="identity",
store_values_multiplied_by_n=False,
)

self.assertIsInstance(converted, _CircularFourierSubclass)

def test_function_value_factory_preserves_requested_subclass(self):
function_values = array([1.0, 0.9, 0.8, 0.7, 0.6, 0.7, 0.8, 0.9, 1.0])

converted = _CircularFourierSubclass.from_function_values(
function_values,
transformation="identity",
store_values_multiplied_by_n=False,
)

self.assertIsInstance(converted, _CircularFourierSubclass)


if __name__ == "__main__":
unittest.main()
Loading