diff --git a/src/pyrecest/filters/hyperhemispherical_grid_filter.py b/src/pyrecest/filters/hyperhemispherical_grid_filter.py index 44ee345f6..9f45ad215 100644 --- a/src/pyrecest/filters/hyperhemispherical_grid_filter.py +++ b/src/pyrecest/filters/hyperhemispherical_grid_filter.py @@ -196,7 +196,8 @@ def update_identity(self, meas_noise, z): Supported noise: :class:`HyperhemisphericalWatsonDistribution`, :class:`WatsonDistribution`, :class:`VonMisesFisherDistribution` - for measurements on the equator, symmetric :class:`HypersphericalMixture`. + for measurements on the equator, symmetric two-component + :class:`HypersphericalMixture` of VMF distributions. Parameters ---------- @@ -233,9 +234,22 @@ def update_identity(self, meas_noise, z): isinstance(meas_noise, HypersphericalMixture) and len(meas_noise.dists) == 2 and all(abs(w - 0.5) < 1e-12 for w in meas_noise.w) + and all( + isinstance(dist, VonMisesFisherDistribution) + for dist in meas_noise.dists + ) + and allclose( + meas_noise.dists[0].mu, -meas_noise.dists[1].mu, atol=1e-12 + ) + and meas_noise.dists[0].kappa == meas_noise.dists[1].kappa ): - meas_noise.dists[0].mu = z - meas_noise.dists[1].mu = -z + meas_noise = HypersphericalMixture( + [ + meas_noise.dists[0].set_mode(z), + meas_noise.dists[1].set_mode(-z), + ], + meas_noise.w, + ) else: raise ValueError( "UpdateIdentity:UnsupportedNoise: unsupported measurement noise type." @@ -283,7 +297,8 @@ def sys_noise_to_transition_density(d_sys, no_grid_points): ---------- d_sys : AbstractDistribution Supported: :class:`HyperhemisphericalWatsonDistribution`, - :class:`WatsonDistribution`, symmetric :class:`HypersphericalMixture`. + :class:`WatsonDistribution`, symmetric two-component + :class:`HypersphericalMixture` of VMF distributions. no_grid_points : int Number of grid points on the hemisphere. @@ -309,6 +324,10 @@ def trans_cp(grid, _grid): isinstance(d_sys, HypersphericalMixture) and len(d_sys.dists) == 2 and all(abs(w - 0.5) < 1e-12 for w in d_sys.w) + and all( + isinstance(dist, VonMisesFisherDistribution) + for dist in d_sys.dists + ) and allclose(d_sys.dists[0].mu, -d_sys.dists[1].mu, atol=1e-12) and d_sys.dists[0].kappa == d_sys.dists[1].kappa ): diff --git a/tests/filters/test_hyperhemispherical_grid_filter_mixture_update.py b/tests/filters/test_hyperhemispherical_grid_filter_mixture_update.py new file mode 100644 index 000000000..006cd101c --- /dev/null +++ b/tests/filters/test_hyperhemispherical_grid_filter_mixture_update.py @@ -0,0 +1,93 @@ +import unittest + +import pyrecest.backend +from pyrecest.backend import allclose, array +from pyrecest.distributions.hypersphere_subset.hyperspherical_mixture import ( + HypersphericalMixture, +) +from pyrecest.distributions.hypersphere_subset.von_mises_fisher_distribution import ( + VonMisesFisherDistribution, +) +from pyrecest.distributions.hypersphere_subset.watson_distribution import ( + WatsonDistribution, +) +from pyrecest.filters.hyperhemispherical_grid_filter import ( + HyperhemisphericalGridFilter, +) + + +@unittest.skipIf( + pyrecest.backend.__backend_name__ == "jax", # pylint: disable=no-member + reason="HyperhemisphericalGridFilter is not supported on the JAX backend", +) +class TestHyperhemisphericalGridFilterMixtureUpdate(unittest.TestCase): + def setUp(self): + self.pole = array([0.0, 0.0, 1.0]) + self.measurement = array([1.0, 0.0, 0.0]) + + def _symmetric_vmf_mixture(self, kappa=3.0): + return HypersphericalMixture( + [ + VonMisesFisherDistribution(self.pole, kappa), + VonMisesFisherDistribution(-self.pole, kappa), + ], + array([0.5, 0.5]), + ) + + def test_update_identity_does_not_mutate_measurement_noise(self): + filt = HyperhemisphericalGridFilter(20, 2) + meas_noise = self._symmetric_vmf_mixture() + + filt.update_identity(meas_noise, self.measurement) + + self.assertTrue(bool(allclose(meas_noise.dists[0].mu, self.pole))) + self.assertTrue(bool(allclose(meas_noise.dists[1].mu, -self.pole))) + + def test_update_identity_rejects_unequal_vmf_concentrations(self): + filt = HyperhemisphericalGridFilter(20, 2) + meas_noise = HypersphericalMixture( + [ + VonMisesFisherDistribution(self.pole, 2.0), + VonMisesFisherDistribution(-self.pole, 3.0), + ], + array([0.5, 0.5]), + ) + + with self.assertRaisesRegex(ValueError, "UnsupportedNoise"): + filt.update_identity(meas_noise, self.measurement) + + def test_update_identity_rejects_nonantipodal_vmf_components(self): + filt = HyperhemisphericalGridFilter(20, 2) + meas_noise = HypersphericalMixture( + [ + VonMisesFisherDistribution(self.pole, 3.0), + VonMisesFisherDistribution(array([1.0, 0.0, 0.0]), 3.0), + ], + array([0.5, 0.5]), + ) + + with self.assertRaisesRegex(ValueError, "UnsupportedNoise"): + filt.update_identity(meas_noise, self.measurement) + + def test_update_identity_rejects_nonunit_measurement_for_mixture(self): + filt = HyperhemisphericalGridFilter(20, 2) + meas_noise = self._symmetric_vmf_mixture() + + with self.assertRaisesRegex(ValueError, "normalized"): + filt.update_identity(meas_noise, array([2.0, 0.0, 0.0])) + + def test_prediction_rejects_non_vmf_symmetric_mixture(self): + d_sys = HypersphericalMixture( + [ + WatsonDistribution(self.pole, 3.0), + WatsonDistribution(-self.pole, 3.0), + ], + array([0.5, 0.5]), + ) + + with self.assertRaisesRegex(ValueError, "unsupported distribution"): + HyperhemisphericalGridFilter.sys_noise_to_transition_density(d_sys, 20) + + +if __name__ == "__main__": + unittest.main()