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
27 changes: 23 additions & 4 deletions src/pyrecest/filters/hyperhemispherical_grid_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
----------
Expand Down Expand Up @@ -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."
Expand Down Expand Up @@ -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.

Expand All @@ -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
):
Expand Down
Original file line number Diff line number Diff line change
@@ -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()
Loading