diff --git a/src/pyrecest/filters/interacting_multiple_model_filter.py b/src/pyrecest/filters/interacting_multiple_model_filter.py index 4ed276933b..0e81c4fb53 100644 --- a/src/pyrecest/filters/interacting_multiple_model_filter.py +++ b/src/pyrecest/filters/interacting_multiple_model_filter.py @@ -471,17 +471,22 @@ def _prepare_transition_matrix(transition_matrix, n_models): if pyrecest.backend.any(transition_matrix < 0.0): raise ValueError("transition_matrix must be elementwise nonnegative.") - row_sums = transition_matrix.sum(axis=1) - if pyrecest.backend.any(row_sums <= 0.0): + row_scales = pyrecest.backend.max(transition_matrix, axis=1) + if pyrecest.backend.any(row_scales <= 0.0): raise ValueError( "Each row of transition_matrix must sum to a positive value." ) - if not allclose(row_sums, 1.0): + + rows_are_normalized = bool( + pyrecest.backend.all(transition_matrix <= 1.0) + ) and allclose(transition_matrix.sum(axis=1), 1.0) + if not rows_are_normalized: warnings.warn( "Rows of transition_matrix do not sum to one. Renormalizing rows.", UserWarning, ) - transition_matrix = transition_matrix / row_sums[:, None] + scaled_rows = transition_matrix / row_scales[:, None] + transition_matrix = scaled_rows / scaled_rows.sum(axis=1)[:, None] return transition_matrix @staticmethod diff --git a/tests/filters/test_imm_transition_matrix_overflow.py b/tests/filters/test_imm_transition_matrix_overflow.py new file mode 100644 index 0000000000..79ed5c798e --- /dev/null +++ b/tests/filters/test_imm_transition_matrix_overflow.py @@ -0,0 +1,42 @@ +import unittest + +import numpy as np +import numpy.testing as npt + +from pyrecest.filters.interacting_multiple_model_filter import ( + InteractingMultipleModelFilter, +) + + +class TestIMMTransitionMatrixOverflow(unittest.TestCase): + def test_normalizes_finite_rows_without_overflow(self): + max_float = np.finfo(float).max + transition_matrix = np.array( + [ + [max_float, max_float / 2.0], + [1.0, 1.0], + ] + ) + + with np.errstate(over="raise", invalid="raise", divide="raise"): + with self.assertWarnsRegex(UserWarning, "Renormalizing rows"): + normalized = ( + InteractingMultipleModelFilter._prepare_transition_matrix( + transition_matrix, 2 + ) + ) + + npt.assert_allclose( + normalized, + np.array( + [ + [2.0 / 3.0, 1.0 / 3.0], + [0.5, 0.5], + ] + ), + ) + npt.assert_allclose(normalized.sum(axis=1), np.ones(2)) + + +if __name__ == "__main__": + unittest.main()