From f41712967237d6cc9309671caec5c4d7094b95b0 Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:31:36 +0300 Subject: [PATCH 1/8] Add binary curve color helper --- .../processing/binary_color_values.py | 56 +++++++++++++++++++ 1 file changed, 56 insertions(+) create mode 100644 src/rtichoke/processing/binary_color_values.py diff --git a/src/rtichoke/processing/binary_color_values.py b/src/rtichoke/processing/binary_color_values.py new file mode 100644 index 00000000..0630c7f4 --- /dev/null +++ b/src/rtichoke/processing/binary_color_values.py @@ -0,0 +1,56 @@ +"""Helpers for applying custom palettes to binary Plotly curves.""" + +from plotly.graph_objs._figure import Figure + + +def _apply_color_values_binary(fig: Figure, color_values) -> Figure: + """Apply custom colors to multiple binary reference groups. + + rtichoke's R implementation keeps a single model black and uses + ``color_values`` only when multiple models or populations are shown. + """ + if color_values is None: + return fig + + reference_groups = [] + for trace in fig.data: + if trace.showlegend is True and trace.name is not None: + name = str(trace.name) + if name not in reference_groups: + reference_groups.append(name) + + # A single model intentionally remains black, matching R. + if len(reference_groups) <= 1: + return fig + + if len(color_values) < len(reference_groups): + raise ValueError( + "color_values must contain at least one color per reference group" + ) + + group_colors = dict(zip(reference_groups, color_values)) + + for trace in fig.data: + legendgroup = "" if trace.legendgroup is None else str(trace.legendgroup) + group = next( + ( + reference_group + for reference_group in reference_groups + if legendgroup == reference_group + or legendgroup.endswith(f"_{reference_group}") + ), + None, + ) + if group is None: + continue + + color = group_colors[group] + if trace.line is not None: + trace.line.color = color + if trace.marker is not None and trace.mode and "markers" in trace.mode: + trace.marker.color = color + if trace.hoverlabel is not None: + trace.hoverlabel.bgcolor = color + trace.hoverlabel.bordercolor = color + + return fig From a0435f8da78cbb27e0c97c79550a3d29b1c30d17 Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:31:57 +0300 Subject: [PATCH 2/8] Honor custom colors in binary ROC curves --- src/rtichoke/discrimination/roc.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/rtichoke/discrimination/roc.py b/src/rtichoke/discrimination/roc.py index 11ce45d8..9884316f 100644 --- a/src/rtichoke/discrimination/roc.py +++ b/src/rtichoke/discrimination/roc.py @@ -4,6 +4,7 @@ from typing import Dict, List, Union, Sequence from plotly.graph_objs._figure import Figure +from rtichoke.processing.binary_color_values import _apply_color_values_binary from rtichoke.processing.plotly_helper_functions import ( _create_rtichoke_plotly_curve_binary, _plot_rtichoke_curve_binary, @@ -87,7 +88,7 @@ def create_roc_curve( color_values=color_values, curve="roc", ) - return fig + return _apply_color_values_binary(fig, color_values) def plot_roc_curve( @@ -210,4 +211,4 @@ def create_roc_curve_times( curve="roc", ) - return fig \ No newline at end of file + return fig From 4ebf45a13f5fe4d6f7eb0aa254126268611f947d Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:32:17 +0300 Subject: [PATCH 3/8] Honor custom colors in binary lift curves --- src/rtichoke/discrimination/lift.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/rtichoke/discrimination/lift.py b/src/rtichoke/discrimination/lift.py index 4efc3256..b0fc7c2d 100644 --- a/src/rtichoke/discrimination/lift.py +++ b/src/rtichoke/discrimination/lift.py @@ -4,6 +4,7 @@ from typing import Dict, List, Sequence, Union from plotly.graph_objs._figure import Figure +from rtichoke.processing.binary_color_values import _apply_color_values_binary from rtichoke.processing.plotly_helper_functions import ( _create_rtichoke_plotly_curve_binary, _plot_rtichoke_curve_binary, @@ -81,7 +82,7 @@ def create_lift_curve( color_values=color_values, curve="lift", ) - return fig + return _apply_color_values_binary(fig, color_values) def plot_lift_curve( From ab10651965f2f5c46426e6d5577b4f4f51b31d9b Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:32:38 +0300 Subject: [PATCH 4/8] Honor custom colors in binary precision-recall curves --- src/rtichoke/discrimination/precision_recall.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/rtichoke/discrimination/precision_recall.py b/src/rtichoke/discrimination/precision_recall.py index 83ccf695..f4af5fab 100644 --- a/src/rtichoke/discrimination/precision_recall.py +++ b/src/rtichoke/discrimination/precision_recall.py @@ -4,6 +4,7 @@ from typing import Dict, List, Sequence, Union from plotly.graph_objs._figure import Figure +from rtichoke.processing.binary_color_values import _apply_color_values_binary from rtichoke.processing.plotly_helper_functions import ( _create_rtichoke_plotly_curve_binary, _plot_rtichoke_curve_binary, @@ -82,7 +83,7 @@ def create_precision_recall_curve( color_values=color_values, curve="precision recall", ) - return fig + return _apply_color_values_binary(fig, color_values) def plot_precision_recall_curve( From 7607c441aa466ce977654abb129c090fe0fc9e7f Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:33:05 +0300 Subject: [PATCH 5/8] Honor custom colors in binary gains curves --- src/rtichoke/discrimination/gains.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/rtichoke/discrimination/gains.py b/src/rtichoke/discrimination/gains.py index 817b88ee..5350446f 100644 --- a/src/rtichoke/discrimination/gains.py +++ b/src/rtichoke/discrimination/gains.py @@ -4,6 +4,7 @@ from typing import Dict, List, Sequence, Union from plotly.graph_objs._figure import Figure +from rtichoke.processing.binary_color_values import _apply_color_values_binary from rtichoke.processing.plotly_helper_functions import ( _create_rtichoke_plotly_curve_times, _create_rtichoke_plotly_curve_binary, @@ -129,7 +130,7 @@ def create_gains_curve( color_values=color_values, curve="gains", ) - return fig + return _apply_color_values_binary(fig, color_values) def plot_gains_curve( From e030709cad14bb96d8c23d96c4e50df14d5ebfda Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:33:33 +0300 Subject: [PATCH 6/8] Honor custom colors in binary decision curves --- src/rtichoke/utility/decision.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/rtichoke/utility/decision.py b/src/rtichoke/utility/decision.py index ab218b74..b8c553f2 100644 --- a/src/rtichoke/utility/decision.py +++ b/src/rtichoke/utility/decision.py @@ -4,6 +4,7 @@ from typing import Dict, List, Sequence, Union from plotly.graph_objs._figure import Figure +from rtichoke.processing.binary_color_values import _apply_color_values_binary from rtichoke.processing.plotly_helper_functions import ( _create_rtichoke_plotly_curve_binary, _plot_rtichoke_curve_binary, @@ -100,7 +101,7 @@ def create_decision_curve( min_p_threshold=min_p_threshold, max_p_threshold=max_p_threshold, ) - return fig + return _apply_color_values_binary(fig, color_values) def plot_decision_curve( @@ -217,7 +218,7 @@ def create_decision_curve_times( max_p_threshold : float, optional The maximum probability threshold to plot. Defaults to 1. by : float, optional - The step size for the probability thresholds. Defaults to 0.01. + The step size for probability thresholds. Defaults to 0.01. stratified_by : Sequence[str], optional Variables for stratification. Defaults to ``["probability_threshold"]``. size : int, optional From f15b0559b5d92560e81bed44c1ba9ac1f431ed77 Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:33:55 +0300 Subject: [PATCH 7/8] Test binary custom color propagation --- tests/test_binary_curve_colors.py | 63 +++++++++++++++++++++++++++++++ 1 file changed, 63 insertions(+) create mode 100644 tests/test_binary_curve_colors.py diff --git a/tests/test_binary_curve_colors.py b/tests/test_binary_curve_colors.py new file mode 100644 index 00000000..e8be33df --- /dev/null +++ b/tests/test_binary_curve_colors.py @@ -0,0 +1,63 @@ +import numpy as np +import pytest + +from rtichoke.discrimination.gains import create_gains_curve +from rtichoke.discrimination.lift import create_lift_curve +from rtichoke.discrimination.precision_recall import create_precision_recall_curve +from rtichoke.discrimination.roc import create_roc_curve +from rtichoke.utility.decision import create_decision_curve + + +PROBS = { + "model_a": np.array([0.1, 0.3, 0.7, 0.9]), + "model_b": np.array([0.2, 0.4, 0.6, 0.8]), +} +REALS = np.array([0, 0, 1, 1]) +CUSTOM_COLORS = ["#111111", "#222222"] + + +@pytest.mark.parametrize( + "creator", + [ + create_roc_curve, + create_precision_recall_curve, + create_lift_curve, + create_gains_curve, + create_decision_curve, + ], +) +def test_binary_create_curves_honor_custom_colors(creator): + fig = creator( + probs=PROBS, + reals=REALS, + by=0.25, + color_values=CUSTOM_COLORS, + ) + + model_traces = [trace for trace in fig.data if trace.showlegend is True] + + assert len(model_traces) == 2 + assert [trace.line.color for trace in model_traces] == CUSTOM_COLORS + + +def test_binary_single_model_remains_black_like_r(): + fig = create_roc_curve( + probs={"model": PROBS["model_a"]}, + reals=REALS, + by=0.25, + color_values=["#123456"], + ) + + model_trace = next(trace for trace in fig.data if trace.name == "model") + + assert model_trace.line.color == "#000000" + + +def test_binary_custom_colors_require_one_per_reference_group(): + with pytest.raises(ValueError, match="one color per reference group"): + create_roc_curve( + probs=PROBS, + reals=REALS, + by=0.25, + color_values=["#111111"], + ) From 0867777801b7d91e4322565d00a3f9eecbc4358a Mon Sep 17 00:00:00 2001 From: Uriah Finkel Date: Thu, 20 Aug 2026 13:34:49 +0300 Subject: [PATCH 8/8] Keep decision documentation unchanged --- src/rtichoke/utility/decision.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/rtichoke/utility/decision.py b/src/rtichoke/utility/decision.py index b8c553f2..7c3c0615 100644 --- a/src/rtichoke/utility/decision.py +++ b/src/rtichoke/utility/decision.py @@ -218,7 +218,7 @@ def create_decision_curve_times( max_p_threshold : float, optional The maximum probability threshold to plot. Defaults to 1. by : float, optional - The step size for probability thresholds. Defaults to 0.01. + The step size for the probability thresholds. Defaults to 0.01. stratified_by : Sequence[str], optional Variables for stratification. Defaults to ``["probability_threshold"]``. size : int, optional