From 27a0fedbb3402b89de45b10db7eeaa248448ad35 Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 14:38:20 +0300 Subject: [PATCH 1/2] Validate and align binary outcome dictionaries --- .../performance_data/performance_data.py | 42 +++++++++++++++++++ 1 file changed, 42 insertions(+) diff --git a/src/rtichoke/performance_data/performance_data.py b/src/rtichoke/performance_data/performance_data.py index 27bac46f..da1efdba 100644 --- a/src/rtichoke/performance_data/performance_data.py +++ b/src/rtichoke/performance_data/performance_data.py @@ -44,6 +44,47 @@ ] +def _validate_and_align_binary_inputs( + probs: Dict[str, np.ndarray], + reals: Union[np.ndarray, Dict[str, np.ndarray]], +) -> Union[np.ndarray, Dict[str, np.ndarray]]: + """Validate binary input alignment and normalize outcome dictionaries.""" + if not isinstance(probs, dict) or not probs: + raise ValueError("`probs` must be a non-empty dictionary of probability arrays.") + + groups = list(probs) + + if isinstance(reals, dict): + expected_keys = set(groups) + if set(reals) != expected_keys: + raise ValueError("`reals` dictionary keys must exactly match `probs`.") + + for group in groups: + n_probs = len(np.asarray(probs[group])) + n_reals = len(np.asarray(reals[group])) + 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 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)) + for group in groups: + n_probs = len(np.asarray(probs[group])) + if n_probs != n_reals: + raise ValueError( + f"Shared outcome length must match probabilities for group {group!r}: " + f"len(probs)={n_probs}, len(reals)={n_reals}." + ) + + return reals + + def _probs_with_r_binary_cutoff_semantics( probs: Dict[str, np.ndarray], by: float ) -> Dict[str, np.ndarray]: @@ -113,6 +154,7 @@ def prepare_binned_classification_data( any other stratification variables. It forms the basis for subsequent performance calculations. """ + reals = _validate_and_align_binary_inputs(probs=probs, reals=reals) breaks = create_breaks_values(None, "probability_threshold", by) aj_data_combinations = _create_aj_data_combinations_binary( From e86479e3ee30f8f78a8dd1265835aaaf29cba1e7 Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 14:38:35 +0300 Subject: [PATCH 2/2] Add binary input alignment regressions --- tests/test_binary_input_validation.py | 69 +++++++++++++++++++++++++++ 1 file changed, 69 insertions(+) create mode 100644 tests/test_binary_input_validation.py diff --git a/tests/test_binary_input_validation.py b/tests/test_binary_input_validation.py new file mode 100644 index 00000000..9b945845 --- /dev/null +++ b/tests/test_binary_input_validation.py @@ -0,0 +1,69 @@ +import numpy as np +import pytest +from polars.testing import assert_frame_equal + +from rtichoke import prepare_performance_data + + +def test_binary_reals_dict_is_aligned_by_group_key_not_insertion_order(): + probs = { + "group_a": np.array([0.1, 0.8, 0.7, 0.2]), + "group_b": np.array([0.9, 0.2, 0.3, 0.8]), + } + reals_a = np.array([0, 1, 1, 0]) + reals_b = np.array([1, 0, 0, 1]) + + aligned = prepare_performance_data( + probs=probs, + reals={"group_a": reals_a, "group_b": reals_b}, + by=0.5, + ) + reversed_order = prepare_performance_data( + probs=probs, + reals={"group_b": reals_b, "group_a": reals_a}, + by=0.5, + ) + + assert_frame_equal(aligned, reversed_order) + + +def test_binary_single_group_reals_dict_matches_array_input(): + probs = {"group_a": np.array([0.1, 0.8, 0.7, 0.2])} + reals = np.array([0, 1, 1, 0]) + + from_array = prepare_performance_data(probs=probs, reals=reals, by=0.5) + from_dict = prepare_performance_data( + probs=probs, + reals={"group_a": reals}, + by=0.5, + ) + + assert_frame_equal(from_array, from_dict) + + +def test_binary_reals_dict_keys_must_match_probability_groups(): + probs = { + "group_a": np.array([0.1, 0.8]), + "group_b": np.array([0.9, 0.2]), + } + + with pytest.raises(ValueError, match="keys must exactly match"): + prepare_performance_data( + probs=probs, + reals={ + "group_a": np.array([0, 1]), + "wrong_group": np.array([1, 0]), + }, + by=0.5, + ) + + +def test_binary_input_lengths_must_match_within_group(): + probs = {"group_a": np.array([0.1, 0.8, 0.7])} + + with pytest.raises(ValueError, match="Input lengths must match"): + prepare_performance_data( + probs=probs, + reals={"group_a": np.array([0, 1])}, + by=0.5, + )