diff --git a/src/pyrecest/utils/pairwise_covariance_features.py b/src/pyrecest/utils/pairwise_covariance_features.py index 286532f9e..211c79b12 100644 --- a/src/pyrecest/utils/pairwise_covariance_features.py +++ b/src/pyrecest/utils/pairwise_covariance_features.py @@ -252,8 +252,13 @@ def _validate_covariance_stack(name: str, covariances: Any) -> Any: def _symmetrized_covariance_batch(covariances: Any) -> Any: + """Move and symmetrize a covariance stack without finite overflow.""" moved = moveaxis(covariances, -1, 0) - return 0.5 * (moved + transpose(moved, (0, 2, 1))) + transposed = transpose(moved, (0, 2, 1)) + scale = maximum(abs(moved), abs(transposed)) + safe_scale = where(scale == 0.0, 1.0, scale) + normalized_average = 0.5 * (moved / safe_scale + transposed / safe_scale) + return scale * normalized_average def _batch_trace(matrices: Any) -> Any: diff --git a/tests/test_pairwise_covariance_symmetrization_overflow.py b/tests/test_pairwise_covariance_symmetrization_overflow.py new file mode 100644 index 000000000..60e345526 --- /dev/null +++ b/tests/test_pairwise_covariance_symmetrization_overflow.py @@ -0,0 +1,34 @@ +import math + +import numpy as np +import numpy.testing as npt + +from pyrecest.backend import array +from pyrecest.utils import pairwise_covariance_shape_components + + +def test_shape_components_preserve_extreme_finite_covariances(): + largest = np.finfo(np.float64).max + covariance_along_first_axis = array( + [ + [[largest], [0.0]], + [[0.0], [0.0]], + ] + ) + covariance_along_second_axis = array( + [ + [[0.0], [0.0]], + [[0.0], [largest]], + ] + ) + + shape_cost, logdet_cost, shape_similarity = ( + pairwise_covariance_shape_components( + covariance_along_first_axis, + covariance_along_second_axis, + ) + ) + + npt.assert_allclose(shape_cost, array([[1.0]])) + npt.assert_allclose(logdet_cost, array([[0.0]])) + npt.assert_allclose(shape_similarity, array([[math.exp(-1.0)]]))