From 785e3f3c95c28cc8c2cceeb4d280dbf5d55b1fb6 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Thu, 6 Aug 2026 17:53:25 +0800 Subject: [PATCH 1/2] Validate hypercylindrical particle filter counts --- .../filters/hypercylindrical_particle_filter.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/src/pyrecest/filters/hypercylindrical_particle_filter.py b/src/pyrecest/filters/hypercylindrical_particle_filter.py index b58fb83d29..8cea3159a9 100644 --- a/src/pyrecest/filters/hypercylindrical_particle_filter.py +++ b/src/pyrecest/filters/hypercylindrical_particle_filter.py @@ -3,11 +3,15 @@ # pylint: disable=redefined-builtin,no-name-in-module,no-member from pyrecest.backend import concatenate, int32, int64, mod, ones, pi, zeros +from pyrecest.distributions.cart_prod.abstract_lin_bounded_cart_prod_distribution import ( + _validate_nonnegative_dimension_count, +) from pyrecest.distributions.cart_prod.hypercylindrical_dirac_distribution import ( HypercylindricalDiracDistribution, ) from .abstract_particle_filter import AbstractParticleFilter +from .hypertoroidal_particle_filter import _validate_positive_integer from .manifold_mixins import HypercylindricalFilterMixin @@ -20,6 +24,12 @@ def __init__( bound_dim: Union[int, int32, int64], lin_dim: Union[int, int32, int64], ): + n_particles = _validate_positive_integer(n_particles, "n_particles") + bound_dim = _validate_nonnegative_dimension_count(bound_dim, "bound_dim") + lin_dim = _validate_nonnegative_dimension_count(lin_dim, "lin_dim") + if bound_dim + lin_dim == 0: + raise ValueError("total dimension must be positive") + d = zeros((n_particles, bound_dim + lin_dim)) w = ones(n_particles) / n_particles filter_state = HypercylindricalDiracDistribution(bound_dim, d, w) From 5c1eb76bef5d864caab7858d22341d0c7a2c429d Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Thu, 6 Aug 2026 17:53:46 +0800 Subject: [PATCH 2/2] Test hypercylindrical particle count validation --- .../test_hypercylindrical_particle_filter.py | 48 +++++++++++++++++++ 1 file changed, 48 insertions(+) diff --git a/tests/filters/test_hypercylindrical_particle_filter.py b/tests/filters/test_hypercylindrical_particle_filter.py index dab6400239..003e266b74 100644 --- a/tests/filters/test_hypercylindrical_particle_filter.py +++ b/tests/filters/test_hypercylindrical_particle_filter.py @@ -1,5 +1,6 @@ import unittest +import numpy as np import numpy.testing as npt import pyrecest.backend @@ -31,6 +32,53 @@ def test_initialization(self): self.assertIsNotNone(hpf.filter_state) self.assertEqual(hpf.filter_state.d.shape, (10, self.bound_dim + self.lin_dim)) + def test_initialization_accepts_scalar_numpy_integer_counts(self): + hpf = HypercylindricalParticleFilter( + np.int64(4), np.array(1, dtype=np.int64), np.int64(2) + ) + + self.assertEqual(hpf.filter_state.d.shape, (4, 3)) + self.assertEqual(hpf.filter_state.bound_dim, 1) + self.assertEqual(hpf.filter_state.lin_dim, 2) + + def test_initialization_rejects_invalid_particle_counts(self): + invalid_counts = ( + 0, + -1, + 1.5, + True, + np.bool_(True), + np.array(True), + np.array([1]), + ) + + for n_particles in invalid_counts: + with self.subTest(n_particles=n_particles): + with self.assertRaisesRegex(ValueError, "positive integer"): + HypercylindricalParticleFilter( + n_particles, self.bound_dim, self.lin_dim + ) + + def test_initialization_rejects_invalid_dimension_counts(self): + invalid_bound_dims = (True, np.bool_(True), np.array(True), -1, 1.5, [1]) + for bound_dim in invalid_bound_dims: + with self.subTest(bound_dim=bound_dim): + with self.assertRaisesRegex( + ValueError, "bound_dim must be a non-negative integer" + ): + HypercylindricalParticleFilter(4, bound_dim, self.lin_dim) + + invalid_lin_dims = (False, np.bool_(False), np.array(False), -1, 1.5, [2]) + for lin_dim in invalid_lin_dims: + with self.subTest(lin_dim=lin_dim): + with self.assertRaisesRegex( + ValueError, "lin_dim must be a non-negative integer" + ): + HypercylindricalParticleFilter(4, self.bound_dim, lin_dim) + + with self.assertRaisesRegex(ValueError, "total dimension must be positive"): + HypercylindricalParticleFilter(4, 0, 0) + @unittest.skipIf( pyrecest.backend.__backend_name__ == "jax", reason="Backend not supported" )