Skip to content
36 changes: 35 additions & 1 deletion src/rtichoke/performance_data/performance_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,34 @@
import numpy as np


def _probs_with_r_binary_cutoff_semantics(
probs: Dict[str, np.ndarray], by: float
) -> Dict[str, np.ndarray]:
"""Place exact binary cutoffs on R's predicted-negative side.

R's binary implementation uses ``prob > cutoff``. The shared Python binning
machinery is left-closed because the time-dependent path intentionally
follows a ``prob >= cutoff`` convention. Move only binary probabilities
exactly equal to a public cutoff one representable float downward before
bin assignment; all other probabilities and the public cutoff grid remain
unchanged.
"""
cutoffs = create_breaks_values(None, "probability_threshold", by)
nonzero_cutoffs = cutoffs[cutoffs > 0]
adjusted = {}

for reference_group, values in probs.items():
values_array = np.asarray(values, dtype=float).copy()
for cutoff in nonzero_cutoffs:
equal_to_cutoff = values_array == cutoff
values_array[equal_to_cutoff] = np.nextafter(
values_array[equal_to_cutoff], -np.inf
)
adjusted[reference_group] = values_array

return adjusted


def prepare_binned_classification_data(
probs: Dict[str, np.ndarray],
reals: Union[np.ndarray, Dict[str, np.ndarray]],
Expand Down Expand Up @@ -69,9 +97,15 @@ def prepare_binned_classification_data(
breaks=breaks,
)

probs_for_binning = (
_probs_with_r_binary_cutoff_semantics(probs, by)
if "probability_threshold" in stratified_by
else probs
)

list_data_to_adjust = _create_list_data_to_adjust_binary(
aj_data_combinations,
probs,
probs_for_binning,
reals,
stratified_by=stratified_by,
by=by,
Expand Down
97 changes: 97 additions & 0 deletions tests/test_cutoff_boundary_semantics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
import numpy as np
import polars as pl
from polars.testing import assert_frame_equal

from rtichoke import prepare_performance_data, prepare_performance_data_times


EXPECTED_BINARY_COUNTS_AT_HALF = {
"true_positives": 1.0,
"false_positives": 0.0,
"true_negatives": 1.0,
"false_negatives": 1.0,
}


def test_binary_probability_equal_to_cutoff_is_predicted_negative_like_r():
result = prepare_performance_data(
probs={"model": np.array([0.4, 0.5, 0.6])},
reals=np.array([0, 1, 1]),
by=0.5,
)

row = result.filter(pl.col("chosen_cutoff") == 0.5).row(0, named=True)

for column, expected in EXPECTED_BINARY_COUNTS_AT_HALF.items():
assert row[column] == expected


def test_time_probability_equal_to_cutoff_remains_positive_for_dcurves_parity():
result = prepare_performance_data_times(
probs={"model": np.array([0.4, 0.5, 0.6])},
reals=np.array([1, 1, 1]),
times=np.array([3.0, 1.0, 1.0]),
fixed_time_horizons=[2.0],
by=0.5,
)

row = result.filter(pl.col("chosen_cutoff") == 0.5).row(0, named=True)

assert row["true_positives"] == 2.0
assert row["false_positives"] == 0.0
assert row["true_negatives"] == 1.0
assert row["false_negatives"] == 0.0


def test_binary_cutoff_zero_still_predicts_everyone_positive():
result = prepare_performance_data(
probs={"model": np.array([0.0, 0.5, 1.0])},
reals=np.array([0, 1, 1]),
by=0.5,
)

row = result.filter(pl.col("chosen_cutoff") == 0.0).row(0, named=True)

assert row["predicted_positives"] == 3
assert row["true_negatives"] == 0
assert row["false_negatives"] == 0


def test_binary_cutoff_one_predicts_everyone_negative():
result = prepare_performance_data(
probs={"model": np.array([0.0, 0.5, 1.0])},
reals=np.array([0, 1, 1]),
by=0.5,
)

row = result.filter(pl.col("chosen_cutoff") == 1.0).row(0, named=True)

assert row["predicted_positives"] == 0
assert row["true_positives"] == 0
assert row["false_positives"] == 0


def test_binary_cutoff_adjustment_does_not_change_ppcr_stratification():
probs = {"model": np.array([0.1, 0.2, 0.5, 0.5, 0.8, 0.9])}
reals = np.array([0, 0, 1, 0, 1, 1])

ppcr_only = prepare_performance_data(
probs=probs,
reals=reals,
stratified_by=("ppcr",),
by=0.5,
).sort(["reference_group", "chosen_cutoff"])

combined_ppcr = (
prepare_performance_data(
probs=probs,
reals=reals,
stratified_by=("probability_threshold", "ppcr"),
by=0.5,
)
.filter(pl.col("stratified_by") == "ppcr")
.sort(["reference_group", "chosen_cutoff"])
.select(ppcr_only.columns)
)

assert_frame_equal(combined_ppcr, ppcr_only)
Loading