From 7975367235df719cd3b3eafd68eab6e2c2af0722 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 13 Jul 2026 15:01:50 +0200 Subject: [PATCH 1/2] Stabilize IMM transition normalization --- .../filters/interacting_multiple_model_filter.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) 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 From f69caeebe75fa6051dc1dfc69cf2161c5ba25fef Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 13 Jul 2026 15:02:05 +0200 Subject: [PATCH 2/2] Add IMM transition overflow regression --- .../test_imm_transition_matrix_overflow.py | 42 +++++++++++++++++++ 1 file changed, 42 insertions(+) create mode 100644 tests/filters/test_imm_transition_matrix_overflow.py 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()