diff --git a/src/pyrecest/tracking/ellipse_geometry.py b/src/pyrecest/tracking/ellipse_geometry.py index 163e55f63..489ca7230 100644 --- a/src/pyrecest/tracking/ellipse_geometry.py +++ b/src/pyrecest/tracking/ellipse_geometry.py @@ -37,7 +37,7 @@ def symmetrize(matrix): """Return the symmetric part of ``matrix``.""" matrix = asarray(matrix) - return 0.5 * (matrix + matrix.T) + return 0.5 * matrix + 0.5 * matrix.T def project_symmetric_covariance(covariance, minimum_eigenvalue=0.0): diff --git a/tests/tracking/test_ellipse_geometry_symmetrization_overflow.py b/tests/tracking/test_ellipse_geometry_symmetrization_overflow.py new file mode 100644 index 000000000..4709ed92c --- /dev/null +++ b/tests/tracking/test_ellipse_geometry_symmetrization_overflow.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +import numpy as np +import numpy.testing as npt + +from pyrecest.tracking.ellipse_geometry import symmetrize + + +def test_ellipse_geometry_symmetrization_avoids_intermediate_overflow() -> None: + matrix = np.full((2, 2), np.finfo(np.float32).max, dtype=np.float32) + + with np.errstate(over="raise", invalid="raise"): + result = symmetrize(matrix) + + result = np.asarray(result) + npt.assert_array_equal(result, matrix) + assert np.all(np.isfinite(result))