diff --git a/src/pyrecest/filters/se2_ukf.py b/src/pyrecest/filters/se2_ukf.py index 5dc7df5b1..41939b5df 100644 --- a/src/pyrecest/filters/se2_ukf.py +++ b/src/pyrecest/filters/se2_ukf.py @@ -47,6 +47,15 @@ def _to_python_bool(value): return bool(value) +def _is_complex_array(value): + """Return whether a NumPy/JAX array or PyTorch tensor has complex dtype.""" + dtype = getattr(value, "dtype", None) + if getattr(dtype, "kind", None) == "c": + return True + is_complex = getattr(value, "is_complex", None) + return bool(is_complex()) if callable(is_complex) else False + + def _normalize_rotation_columns(rotation_samples, fallback_rotation): """Normalize 2-D rotation columns, replacing undefined zero directions.""" rotation_samples = asarray(rotation_samples) @@ -72,6 +81,10 @@ def _validate_se2_gaussian(distribution, role): mu = asarray(distribution.mu) covariance = asarray(distribution.C) + if _is_complex_array(mu): + raise ValueError(f"{role} mean must be real-valued.") + if _is_complex_array(covariance): + raise ValueError(f"{role} covariance must be real-valued.") if mu.shape != (4,): raise ValueError(f"{role} mean must be a 4-D vector.") if covariance.shape != (4, 4): @@ -95,6 +108,8 @@ def _validate_se2_gaussian(distribution, role): def _validate_se2_measurement(z): measurement = asarray(z) + if _is_complex_array(measurement): + raise ValueError("measurement z must be real-valued.") if measurement.shape != (4,): raise ValueError("measurement z must be a 4-D vector.") if not _to_python_bool(backend_all(isfinite(measurement))): diff --git a/tests/filters/test_se2_ukf_real_inputs.py b/tests/filters/test_se2_ukf_real_inputs.py new file mode 100644 index 000000000..ea2e25c61 --- /dev/null +++ b/tests/filters/test_se2_ukf_real_inputs.py @@ -0,0 +1,57 @@ +"""Regression tests for real-valued SE(2) UKF inputs.""" + +import unittest + +import numpy.testing as npt + +# pylint: disable=no-name-in-module,no-member +import pyrecest.backend +from pyrecest.backend import array, eye, to_numpy +from pyrecest.distributions import GaussianDistribution +from pyrecest.filters.se2_ukf import SE2UKF + + +@unittest.skipIf( + pyrecest.backend.__backend_name__ == "jax", + reason="SE2UKF update is not supported on JAX", +) +class TestSE2UKFRealInputs(unittest.TestCase): + @staticmethod + def _noise_distribution(): + return GaussianDistribution( + array([1.0, 0.0, 0.0, 0.0]), + 0.1 * eye(4), + ) + + def test_update_rejects_complex_measurement_without_mutating_state(self): + current_filter = SE2UKF() + original_mean = to_numpy(current_filter.filter_state.mu).copy() + original_covariance = to_numpy(current_filter.filter_state.C).copy() + + with self.assertRaisesRegex(ValueError, "real-valued"): + current_filter.update_identity( + self._noise_distribution(), + array([1.0 + 0.0j, 0.0, 1.0j, 0.0]), + ) + + npt.assert_allclose( + to_numpy(current_filter.filter_state.mu), original_mean + ) + npt.assert_allclose( + to_numpy(current_filter.filter_state.C), original_covariance + ) + + def test_filter_state_rejects_complex_mean_direction(self): + current_filter = SE2UKF() + invalid_state = GaussianDistribution( + array([1.0, 0.0, 0.0, 0.0]), + eye(4), + ) + invalid_state.mu = array([1.0 + 0.0j, 0.0, 1.0j, 0.0]) + + with self.assertRaisesRegex(ValueError, "real-valued"): + current_filter.filter_state = invalid_state + + +if __name__ == "__main__": + unittest.main()