diff --git a/src/pyrecest/filters/hypertoroidal_particle_filter.py b/src/pyrecest/filters/hypertoroidal_particle_filter.py index 2205c4e6b..de97a12e7 100644 --- a/src/pyrecest/filters/hypertoroidal_particle_filter.py +++ b/src/pyrecest/filters/hypertoroidal_particle_filter.py @@ -6,7 +6,6 @@ # pylint: disable=redefined-builtin,no-name-in-module,no-member # pylint: disable=no-name-in-module,no-member from pyrecest.backend import ( - arange, int32, int64, linspace, @@ -65,12 +64,11 @@ def __init__( ): n_particles = _validate_positive_integer(n_particles, "n_particles") dim = _validate_positive_integer(dim, "dim") + points_1d = linspace(0.0, 2.0 * pi, num=n_particles, endpoint=False) if dim == 1: - points = linspace(0.0, 2.0 * pi, num=n_particles, endpoint=False) + points = points_1d else: - points = tile( - arange(0.0, 2.0 * pi, 2.0 * pi / n_particles), (dim, 1) - ).T.squeeze() + points = tile(points_1d, (dim, 1)).T filter_state = HypertoroidalDiracDistribution(points, dim=dim) HypertoroidalFilterMixin.__init__(self) AbstractParticleFilter.__init__(self, filter_state) diff --git a/tests/filters/test_hypertoroidal_particle_filter.py b/tests/filters/test_hypertoroidal_particle_filter.py index d5e2839c6..bf397e1d2 100644 --- a/tests/filters/test_hypertoroidal_particle_filter.py +++ b/tests/filters/test_hypertoroidal_particle_filter.py @@ -41,6 +41,19 @@ def test_constructor_rejects_invalid_dimension(self): with self.assertRaisesRegex(ValueError, "dim"): HypertoroidalParticleFilter(5, dim) + def test_constructor_preserves_requested_particle_count(self): + hpf = HypertoroidalParticleFilter(61, 2) + + self.assertEqual(hpf.filter_state.d.shape, (61, 2)) + self.assertEqual(hpf.filter_state.w.shape, (61,)) + + def test_constructor_preserves_singleton_particle_axis(self): + hpf = HypertoroidalParticleFilter(1, 3) + + self.assertEqual(hpf.filter_state.d.shape, (1, 3)) + self.assertEqual(hpf.filter_state.w.shape, (1,)) + self.assertEqual(hpf.get_point_estimate().shape, (3,)) + @unittest.skipIf( pyrecest.backend.__backend_name__ == "jax", reason="Backend not supported'" )