diff --git a/src/rtichoke/performance_data/performance_data.py b/src/rtichoke/performance_data/performance_data.py index da1efdba..814939a3 100644 --- a/src/rtichoke/performance_data/performance_data.py +++ b/src/rtichoke/performance_data/performance_data.py @@ -53,6 +53,12 @@ def _validate_and_align_binary_inputs( raise ValueError("`probs` must be a non-empty dictionary of probability arrays.") groups = list(probs) + for group in groups: + probs_values = np.asarray(probs[group]) + if not np.all(np.isfinite(probs_values)) or np.any( + (probs_values < 0) | (probs_values > 1) + ): + raise ValueError("Estimated probabilities must be between 0 and 1.") if isinstance(reals, dict): expected_keys = set(groups) @@ -60,20 +66,27 @@ def _validate_and_align_binary_inputs( raise ValueError("`reals` dictionary keys must exactly match `probs`.") for group in groups: + reals_values = np.asarray(reals[group]) n_probs = len(np.asarray(probs[group])) - n_reals = len(np.asarray(reals[group])) + n_reals = len(reals_values) if n_probs != n_reals: raise ValueError( f"Input lengths must match within group {group!r}: " f"len(probs)={n_probs}, len(reals)={n_reals}." ) + if not np.all(np.isin(reals_values, [0, 1])): + raise ValueError("Binary outcomes must contain only 0 and 1.") if len(groups) == 1: return np.asarray(reals[groups[0]]) return {group: np.asarray(reals[group]) for group in groups} - n_reals = len(np.asarray(reals)) + reals_values = np.asarray(reals) + if not np.all(np.isin(reals_values, [0, 1])): + raise ValueError("Binary outcomes must contain only 0 and 1.") + + n_reals = len(reals_values) for group in groups: n_probs = len(np.asarray(probs[group])) if n_probs != n_reals: diff --git a/src/rtichoke/processing/time_input_validation.py b/src/rtichoke/processing/time_input_validation.py index 32fd4c4c..dd1cbe59 100644 --- a/src/rtichoke/processing/time_input_validation.py +++ b/src/rtichoke/processing/time_input_validation.py @@ -5,6 +5,26 @@ import numpy as np +def _validate_probability_values(probs: Dict[str, np.ndarray]) -> None: + for values in probs.values(): + probs_values = np.asarray(values) + if not np.all(np.isfinite(probs_values)) or np.any( + (probs_values < 0) | (probs_values > 1) + ): + raise ValueError("Estimated probabilities must be between 0 and 1.") + + +def _validate_time_outcome_values( + reals: Union[np.ndarray, Dict[str, np.ndarray]], +) -> None: + values = reals.values() if isinstance(reals, dict) else [reals] + for outcome_values in values: + if not np.all(np.isin(np.asarray(outcome_values), [0, 1, 2])): + raise ValueError( + "Time-dependent outcomes must contain only 0, 1, and 2." + ) + + def _validate_time_input_alignment( probs: Dict[str, np.ndarray], reals: Union[np.ndarray, Dict[str, np.ndarray]], @@ -14,6 +34,9 @@ def _validate_time_input_alignment( if not isinstance(probs, dict) or not probs: raise ValueError("`probs` must be a non-empty dictionary of probability arrays.") + _validate_probability_values(probs) + _validate_time_outcome_values(reals) + groups = list(probs) multiple_groups = len(groups) > 1 reals_is_dict = isinstance(reals, dict) diff --git a/tests/test_input_domain_validation.py b/tests/test_input_domain_validation.py new file mode 100644 index 00000000..744176b7 --- /dev/null +++ b/tests/test_input_domain_validation.py @@ -0,0 +1,44 @@ +import numpy as np +import pytest + +from rtichoke import prepare_performance_data, prepare_performance_data_times + + +def test_binary_probabilities_must_be_in_unit_interval(): + with pytest.raises(ValueError, match="between 0 and 1"): + prepare_performance_data( + probs={"model": np.array([0.1, 1.1])}, + reals=np.array([0, 1]), + by=0.5, + ) + + +def test_binary_outcomes_must_be_zero_or_one(): + with pytest.raises(ValueError, match="only 0 and 1"): + prepare_performance_data( + probs={"model": np.array([0.1, 0.9])}, + reals=np.array([0, 2]), + by=0.5, + ) + + +def test_time_probabilities_must_be_in_unit_interval(): + with pytest.raises(ValueError, match="between 0 and 1"): + prepare_performance_data_times( + probs={"model": np.array([-0.1, 0.9])}, + reals=np.array([0, 1]), + times=np.array([1.0, 2.0]), + fixed_time_horizons=[1.5], + by=0.5, + ) + + +def test_time_outcomes_must_use_supported_event_codes(): + with pytest.raises(ValueError, match="only 0, 1, and 2"): + prepare_performance_data_times( + probs={"model": np.array([0.1, 0.9])}, + reals=np.array([0, 3]), + times=np.array([1.0, 2.0]), + fixed_time_horizons=[1.5], + by=0.5, + )