diff --git a/src/pyrecest/numerics.py b/src/pyrecest/numerics.py index 2b28e179d..1bd455372 100644 --- a/src/pyrecest/numerics.py +++ b/src/pyrecest/numerics.py @@ -256,9 +256,14 @@ def nearest_symmetric_psd(matrix, *, min_eigenvalue: float = 0.0): _raise_if_not_square_matrix(arr) _raise_if_nonfinite_matrix(arr, "matrix") sym = _stable_symmetric_average(arr) - eigvals, eigvecs = np.linalg.eigh(sym) - clipped = np.maximum(eigvals, min_eigenvalue) - repaired = (eigvecs * clipped) @ eigvecs.T + if sym.size == 0: + return _from_numpy_array(sym) + + scale = max(1.0, float(np.max(np.abs(sym))), min_eigenvalue) + scaled_sym = sym / scale + eigvals, eigvecs = np.linalg.eigh(scaled_sym) + clipped = np.maximum(eigvals, min_eigenvalue / scale) + repaired = ((eigvecs * clipped) @ eigvecs.T) * scale return _from_numpy_array(_stable_symmetric_average(repaired)) diff --git a/tests/test_numerics_extreme_psd_projection.py b/tests/test_numerics_extreme_psd_projection.py new file mode 100644 index 000000000..96317d1ad --- /dev/null +++ b/tests/test_numerics_extreme_psd_projection.py @@ -0,0 +1,19 @@ +import numpy as np + +from pyrecest.numerics import nearest_symmetric_psd + + +def test_nearest_symmetric_psd_scales_dense_extreme_covariance_before_eigh(): + maximum = np.finfo(float).max + matrix = np.full((2, 2), maximum) + + with np.errstate(over="raise", invalid="raise"): + repaired = np.asarray(nearest_symmetric_psd(matrix)) + + assert np.all(np.isfinite(repaired)) + np.testing.assert_allclose( + repaired, + matrix, + rtol=8.0 * np.finfo(float).eps, + atol=0.0, + )