diff --git a/src/pyrecest/filters/so3_product_particle_filter.py b/src/pyrecest/filters/so3_product_particle_filter.py index d3557eab5..4e2c4eb4a 100644 --- a/src/pyrecest/filters/so3_product_particle_filter.py +++ b/src/pyrecest/filters/so3_product_particle_filter.py @@ -15,7 +15,6 @@ exp, isfinite, isnan, - linalg, log, max, maximum, @@ -164,7 +163,7 @@ def _as_particle_array(particles, num_rotations): flat_particles = reshape(particles, (-1, 4)) if not all(isfinite(flat_particles)): raise ValueError("SO(3)^K particles must be finite.") - if not all(linalg.norm(flat_particles, axis=-1) > 0.0): + if not all(any(flat_particles != 0.0, axis=-1)): raise ValueError("SO(3)^K particles must be nonzero.") normalized = normalize_quaternions(flat_particles) diff --git a/tests/filters/test_so3_product_particle_filter_extreme_quaternions.py b/tests/filters/test_so3_product_particle_filter_extreme_quaternions.py new file mode 100644 index 000000000..3acf31353 --- /dev/null +++ b/tests/filters/test_so3_product_particle_filter_extreme_quaternions.py @@ -0,0 +1,26 @@ +import numpy as np +import numpy.testing as npt + +# pylint: disable=no-name-in-module,no-member +from pyrecest.backend import array, to_numpy +from pyrecest.filters import SO3ProductParticleFilter + + +def test_extreme_finite_quaternions_normalize_without_overflow(): + backend_dtype = to_numpy(array([1.0])).dtype + largest = np.finfo(backend_dtype).max + particles = array([[[largest / 2.0, largest / 2.0, 0.0, 0.0]]]) + + with np.errstate(over="raise", invalid="raise", divide="raise"): + filt = SO3ProductParticleFilter( + n_particles=1, + num_rotations=1, + initial_particles=particles, + ) + + expected = np.array( + [1.0 / np.sqrt(2.0), 1.0 / np.sqrt(2.0), 0.0, 0.0] + ) + actual = to_numpy(filt.particles[0, 0]) + npt.assert_allclose(actual, expected, rtol=1e-6, atol=0.0) + npt.assert_allclose(np.linalg.norm(actual), 1.0, rtol=1e-6)