Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions src/pyrecest/filters/se2_ukf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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):
Expand All @@ -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))):
Expand Down
57 changes: 57 additions & 0 deletions tests/filters/test_se2_ukf_real_inputs.py
Original file line number Diff line number Diff line change
@@ -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()
Loading