From 60348238c2bee63fd5ac4582702335cbb024aa08 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Fri, 31 Jul 2026 10:37:33 +0200 Subject: [PATCH 1/3] Stabilize pairwise covariance symmetrization --- src/pyrecest/utils/pairwise_covariance_features.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/src/pyrecest/utils/pairwise_covariance_features.py b/src/pyrecest/utils/pairwise_covariance_features.py index 286532f9e1..4030131886 100644 --- a/src/pyrecest/utils/pairwise_covariance_features.py +++ b/src/pyrecest/utils/pairwise_covariance_features.py @@ -252,8 +252,15 @@ 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: From b4a6c2c9d46075a23b63b03c1500725f975500f4 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Fri, 31 Jul 2026 10:38:05 +0200 Subject: [PATCH 2/3] Test extreme pairwise covariance scales --- ...wise_covariance_symmetrization_overflow.py | 34 +++++++++++++++++++ 1 file changed, 34 insertions(+) create mode 100644 tests/test_pairwise_covariance_symmetrization_overflow.py diff --git a/tests/test_pairwise_covariance_symmetrization_overflow.py b/tests/test_pairwise_covariance_symmetrization_overflow.py new file mode 100644 index 0000000000..60e345526f --- /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)]])) From c571c3a2f0da0610a459850f544d20d983040455 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Fri, 31 Jul 2026 10:40:06 +0200 Subject: [PATCH 3/3] Format stable covariance average --- src/pyrecest/utils/pairwise_covariance_features.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/pyrecest/utils/pairwise_covariance_features.py b/src/pyrecest/utils/pairwise_covariance_features.py index 4030131886..211c79b12a 100644 --- a/src/pyrecest/utils/pairwise_covariance_features.py +++ b/src/pyrecest/utils/pairwise_covariance_features.py @@ -257,9 +257,7 @@ def _symmetrized_covariance_batch(covariances: Any) -> Any: 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 - ) + normalized_average = 0.5 * (moved / safe_scale + transposed / safe_scale) return scale * normalized_average