From 134789f93b5be688b155e4cf433fe1f9675a47a9 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Sat, 8 Aug 2026 01:00:38 +0800 Subject: [PATCH 1/2] Fix hypertoroidal particle grid shape and count --- src/pyrecest/filters/hypertoroidal_particle_filter.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/pyrecest/filters/hypertoroidal_particle_filter.py b/src/pyrecest/filters/hypertoroidal_particle_filter.py index 2205c4e6bd..de97a12e78 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) From 22538a4e66c63dd544af0d9b1e7d6bf5b7bc7ccc Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Sat, 8 Aug 2026 01:00:57 +0800 Subject: [PATCH 2/2] Test hypertoroidal particle count and singleton axis --- tests/filters/test_hypertoroidal_particle_filter.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/tests/filters/test_hypertoroidal_particle_filter.py b/tests/filters/test_hypertoroidal_particle_filter.py index d5e2839c6d..bf397e1d23 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'" )