Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 4 additions & 16 deletions src/rtichoke/discrimination/gains.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,11 @@
from rtichoke.processing.plotly_helper_functions import (
_create_rtichoke_plotly_curve_binary,
_plot_rtichoke_curve_binary,
_create_rtichoke_curve_list_times,
_create_plotly_curve_times,
_create_reference_lines_data,
_check_if_multiple_populations_are_being_validated_times,
)
from rtichoke.processing.time_reference_lines import _apply_color_values_times
from rtichoke.performance_data.performance_data_times import (
prepare_performance_data_times,
from rtichoke.processing.time_reference_lines import (
_create_rtichoke_plotly_curve_times_reference_safe,
)
import numpy as np
import polars as pl
Expand Down Expand Up @@ -237,24 +234,15 @@ def create_gains_curve_times(
Figure
A Plotly ``Figure`` object for the time-dependent Gains curve.
"""
performance_data = prepare_performance_data_times(
return _create_rtichoke_plotly_curve_times_reference_safe(
probs,
reals,
times,
fixed_time_horizons=fixed_time_horizons,
heuristics_sets=heuristics_sets,
by=by,
stratified_by=stratified_by,
)

curve_list = _create_rtichoke_curve_list_times(
performance_data,
stratified_by=stratified_by[0],
size=size,
color_value=color_values,
color_values=color_values,
curve="gains",
)
curve_list = _apply_color_values_times(curve_list, color_values)
curve_list = _replace_gains_reference_data_times(curve_list, performance_data)

return _create_plotly_curve_times(curve_list)
59 changes: 58 additions & 1 deletion tests/test_time_reference_ownership.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import polars as pl
import pytest

from rtichoke import create_precision_recall_curve_times
from rtichoke import create_gains_curve_times, create_precision_recall_curve_times
from rtichoke.processing.time_reference_lines import _replace_reference_data_times

HORIZONS = [5.0, 10.0]
Expand Down Expand Up @@ -39,6 +39,14 @@ def _nonempty_random_reference_names(fig) -> set[str]:
}


def _nonempty_perfect_reference_names(fig) -> set[str]:
return {
trace.name
for trace in fig.data
if "perfect_model" in trace.name and len(trace.x) > 0
}


@pytest.mark.parametrize(
("curve", "expected"),
[
Expand Down Expand Up @@ -138,3 +146,52 @@ def test_public_multiple_models_still_share_population_reference():
)

assert _nonempty_random_reference_names(fig) == {"random_guess"}


def test_public_equal_risk_populations_keep_distinct_gains_references():
probs = {
"Population A": np.array([0.05, 0.15, 0.35, 0.55, 0.75, 0.95]),
"Population B": np.array([0.10, 0.25, 0.45, 0.65, 0.80, 0.90]),
}
reals = {
"Population A": np.array([1, 1, 0, 0, 0, 0]),
"Population B": np.array([1, 1, 0, 0, 0, 0]),
}
times = {
"Population A": np.array([3.0, 8.0, 12.0, 13.0, 14.0, 15.0]),
"Population B": np.array([4.0, 9.0, 12.0, 13.0, 14.0, 15.0]),
}

fig = create_gains_curve_times(
probs,
reals,
times,
fixed_time_horizons=HORIZONS,
heuristics_sets=HEURISTICS,
by=0.25,
)

assert _nonempty_perfect_reference_names(fig) == {
"perfect_model_Population A",
"perfect_model_Population B",
}


def test_public_gains_models_share_one_population_reference():
probs = {
"Model A": np.array([0.05, 0.15, 0.35, 0.55, 0.75, 0.95]),
"Model B": np.array([0.10, 0.25, 0.45, 0.65, 0.80, 0.90]),
}
reals = np.array([1, 1, 0, 0, 0, 0])
times = np.array([3.0, 8.0, 12.0, 13.0, 14.0, 15.0])

fig = create_gains_curve_times(
probs,
reals,
times,
fixed_time_horizons=HORIZONS,
heuristics_sets=HEURISTICS,
by=0.25,
)

assert _nonempty_perfect_reference_names(fig) == {"perfect_model"}
Loading