diff --git a/src/pyrecest/filters/tracklet_viterbi.py b/src/pyrecest/filters/tracklet_viterbi.py index 9c2318cfa..6d0e0a937 100644 --- a/src/pyrecest/filters/tracklet_viterbi.py +++ b/src/pyrecest/filters/tracklet_viterbi.py @@ -71,6 +71,13 @@ class TrackletAssociationCandidate: velocity: Any | None = None metadata: Mapping[str, Any] = field(default_factory=dict) + def __post_init__(self) -> None: + object.__setattr__( + self, + "unary_cost", + _as_scalar_float(self.unary_cost, "unary_cost"), + ) + @dataclass(frozen=True) class TrackletViterbiConfig: @@ -283,12 +290,13 @@ def _solve_tracklet_viterbi( for previous_miss_streak, previous_cost in state_costs[-1][ previous_index ].items(): - transition_value = float( + transition_value = _as_scalar_float( transition( previous_node.candidate, current_node.candidate, previous_miss_streak, - ) + ), + "transition_cost", ) current_miss_streak = ( previous_miss_streak + 1 if current_node.is_miss else 0 diff --git a/tests/filters/test_tracklet_viterbi.py b/tests/filters/test_tracklet_viterbi.py index 69f6f92d7..028edc4aa 100644 --- a/tests/filters/test_tracklet_viterbi.py +++ b/tests/filters/test_tracklet_viterbi.py @@ -268,3 +268,30 @@ def transition(previous, current, miss_streak): (1,), (2,), ] + + +@pytest.mark.parametrize("invalid_cost", [np.nan, np.inf, -np.inf]) +def test_tracklet_candidate_rejects_nonfinite_unary_costs(invalid_cost): + with pytest.raises(ValueError, match="unary_cost must be finite"): + TrackletAssociationCandidate("invalid", unary_cost=invalid_cost) + + +@pytest.mark.parametrize("invalid_cost", [np.nan, np.inf, -np.inf]) +def test_tracklet_viterbi_solvers_reject_nonfinite_transition_costs(invalid_cost): + frames = [ + [TrackletAssociationCandidate("a", unary_cost=0.0)], + [TrackletAssociationCandidate("b", unary_cost=0.0)], + ] + + def transition(_previous, _current, _miss_streak): + return invalid_cost + + with pytest.raises(ValueError, match="transition_cost must be finite"): + solve_tracklet_viterbi(frames, transition_cost=transition) + + with pytest.raises(ValueError, match="transition_cost must be finite"): + solve_fixed_lag_tracklet_viterbi( + frames, + lag_s=0.1, + transition_cost=transition, + )