diff --git a/src/pyrecest/distributions/hypersphere_subset/abstract_hypersphere_subset_distribution.py b/src/pyrecest/distributions/hypersphere_subset/abstract_hypersphere_subset_distribution.py index 0ae1b395e..28416bcf8 100644 --- a/src/pyrecest/distributions/hypersphere_subset/abstract_hypersphere_subset_distribution.py +++ b/src/pyrecest/distributions/hypersphere_subset/abstract_hypersphere_subset_distribution.py @@ -339,8 +339,10 @@ def hellinger_distance(pdf1, pdf2): ) ) - distance_integral = self.__class__.integrate_fun_over_domain( - fangles_hellinger, self.dim + distance_integral = ( + AbstractHypersphereSubsetDistribution.integrate_fun_over_domain_part( + fangles_hellinger, integration_boundaries + ) ) return sqrt(0.5 * distance_integral) @@ -365,8 +367,10 @@ def total_variation_distance(pdf1, pdf2): ) ) - distance_integral = self.__class__.integrate_fun_over_domain( - fangles_total_variation, self.dim + distance_integral = ( + AbstractHypersphereSubsetDistribution.integrate_fun_over_domain_part( + fangles_total_variation, integration_boundaries + ) ) return 0.5 * distance_integral diff --git a/tests/distributions/test_hypersphere_subset_distance_validation.py b/tests/distributions/test_hypersphere_subset_distance_validation.py index 6e05d0b90..f6685ce00 100644 --- a/tests/distributions/test_hypersphere_subset_distance_validation.py +++ b/tests/distributions/test_hypersphere_subset_distance_validation.py @@ -1,5 +1,8 @@ import unittest +import pyrecest.backend +from pyrecest.backend import array, pi +from pyrecest.distributions import VonMisesFisherDistribution from pyrecest.distributions.hypersphere_subset.hyperspherical_uniform_distribution import ( HypersphericalUniformDistribution, ) @@ -20,6 +23,29 @@ def test_total_variation_distance_rejects_dimension_mismatch(self): with self.assertRaisesRegex(ValueError, "different number of dimensions"): dist.total_variation_distance_numerical(other) + @unittest.skipIf( + pyrecest.backend.__backend_name__ == "jax", + "Numerical hyperspherical integration is not supported on JAX.", + ) + def test_distances_respect_custom_integration_boundaries(self): + dist = VonMisesFisherDistribution(array([1.0, 0.0]), 2.0) + other = VonMisesFisherDistribution(array([0.0, 1.0]), 1.0) + partial_boundaries = array([[0.0, pi / 2.0]]) + + for distance_name in ( + "hellinger_distance_numerical", + "total_variation_distance_numerical", + ): + with self.subTest(distance=distance_name): + distance = getattr(dist, distance_name) + full_distance = float(distance(other)) + partial_distance = float( + distance(other, integration_boundaries=partial_boundaries) + ) + + self.assertGreater(partial_distance, 0.0) + self.assertLess(partial_distance, full_distance) + if __name__ == "__main__": unittest.main()