diff --git a/src/pyrecest/filters/discrete_state/__init__.py b/src/pyrecest/filters/discrete_state/__init__.py index b08fd75b9d..1a36ca1763 100644 --- a/src/pyrecest/filters/discrete_state/__init__.py +++ b/src/pyrecest/filters/discrete_state/__init__.py @@ -114,10 +114,11 @@ def _validated_probability_vector( mask = _module_globals["_coerce_valid_state_mask"](valid_state_mask, n_entries) if mask is not None: values[~mask] = 0.0 - total = float(values.sum()) - if total <= 0.0: + scale = float(values.max()) + if scale <= 0.0: raise ValueError(f"{name} must contain positive probability mass") - return values / total + values /= scale + return values / float(values.sum()) def sparse_gaussian_transition_matrix( diff --git a/tests/filters/test_discrete_state_probability_overflow.py b/tests/filters/test_discrete_state_probability_overflow.py new file mode 100644 index 0000000000..5f93a14b84 --- /dev/null +++ b/tests/filters/test_discrete_state_probability_overflow.py @@ -0,0 +1,25 @@ +import unittest + +import numpy as np +from pyrecest.filters.discrete_state import discrete_forward_backward + + +class TestDiscreteStateProbabilityOverflow(unittest.TestCase): + def test_large_finite_initial_probabilities_keep_relative_mass(self): + largest = np.finfo(float).max + initial_probabilities = np.array([largest, largest / 2.0]) + + result = discrete_forward_backward( + np.zeros((1, 2)), + np.eye(2), + initial_probabilities=initial_probabilities, + ) + + np.testing.assert_allclose( + result.filtered_probabilities[0], + np.array([2.0 / 3.0, 1.0 / 3.0]), + ) + + +if __name__ == "__main__": + unittest.main()