Skip to content

Commit 82dfead

Browse files
authored
Merge pull request #334 from uriahf/fix/color-values-propagation
Honor custom colors in time-dependent curves
2 parents b1453ca + 9be2f21 commit 82dfead

3 files changed

Lines changed: 96 additions & 0 deletions

File tree

src/rtichoke/discrimination/gains.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
_create_reference_lines_data,
1414
_check_if_multiple_populations_are_being_validated_times,
1515
)
16+
from rtichoke.processing.time_reference_lines import _apply_color_values_times
1617
from rtichoke.performance_data.performance_data_times import prepare_performance_data_times
1718
import numpy as np
1819
import polars as pl
@@ -251,6 +252,7 @@ def create_gains_curve_times(
251252
color_value=color_values,
252253
curve="gains",
253254
)
255+
curve_list = _apply_color_values_times(curve_list, color_values)
254256
curve_list = _replace_gains_reference_data_times(curve_list, performance_data)
255257

256258
return _create_plotly_curve_times(curve_list)

src/rtichoke/processing/time_reference_lines.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,33 @@ def _get_reference_aj_estimates_times(performance_data: pl.DataFrame) -> pl.Data
2727
)
2828

2929

30+
def _apply_color_values_times(curve_list: dict, color_values) -> dict:
31+
"""Apply custom colors while retaining R's single-model black styling."""
32+
if color_values is None or not curve_list["multiple_reference_groups"]:
33+
return curve_list
34+
35+
reference_groups = curve_list["reference_group_keys"]
36+
if len(color_values) < len(reference_groups):
37+
raise ValueError(
38+
"color_values must contain at least one color per reference group"
39+
)
40+
41+
colors_dictionary = curve_list["colors_dictionary"]
42+
for index, reference_group in enumerate(reference_groups):
43+
color = color_values[index]
44+
for key in (
45+
reference_group,
46+
f"random_guess_{reference_group}",
47+
f"perfect_model_{reference_group}",
48+
f"treat_none_{reference_group}",
49+
f"treat_all_{reference_group}",
50+
):
51+
if key in colors_dictionary:
52+
colors_dictionary[key] = color
53+
54+
return curve_list
55+
56+
3057
def _replace_reference_data_times(
3158
curve_list: dict,
3259
performance_data: pl.DataFrame,
@@ -88,6 +115,7 @@ def _create_rtichoke_plotly_curve_times_reference_safe(
88115
min_p_threshold=min_p_threshold,
89116
max_p_threshold=max_p_threshold,
90117
)
118+
curve_list = _apply_color_values_times(curve_list, color_values)
91119
return _create_plotly_curve_times(
92120
_replace_reference_data_times(
93121
curve_list,

tests/test_time_curve_colors.py

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,66 @@
1+
import pytest
2+
3+
from rtichoke.processing.time_reference_lines import _apply_color_values_times
4+
5+
6+
def _curve_list(reference_groups, multiple_reference_groups=True):
7+
colors_dictionary = {
8+
"random_guess": "#BEBEBE",
9+
"perfect_model": "#BEBEBE",
10+
"treat_none": "#BEBEBE",
11+
"treat_all": "#BEBEBE",
12+
}
13+
for group in reference_groups:
14+
for key in (
15+
group,
16+
f"random_guess_{group}",
17+
f"perfect_model_{group}",
18+
f"treat_none_{group}",
19+
f"treat_all_{group}",
20+
):
21+
colors_dictionary[key] = "#000000"
22+
23+
return {
24+
"reference_group_keys": reference_groups,
25+
"multiple_reference_groups": multiple_reference_groups,
26+
"colors_dictionary": colors_dictionary,
27+
}
28+
29+
30+
def test_time_curve_custom_colors_propagate_to_groups_and_references():
31+
curve_list = _curve_list(["model_a", "model_b"])
32+
33+
result = _apply_color_values_times(curve_list, ["#111111", "#222222"])
34+
35+
for group, expected_color in (
36+
("model_a", "#111111"),
37+
("model_b", "#222222"),
38+
):
39+
for key in (
40+
group,
41+
f"random_guess_{group}",
42+
f"perfect_model_{group}",
43+
f"treat_none_{group}",
44+
f"treat_all_{group}",
45+
):
46+
assert result["colors_dictionary"][key] == expected_color
47+
48+
assert result["colors_dictionary"]["random_guess"] == "#BEBEBE"
49+
assert result["colors_dictionary"]["perfect_model"] == "#BEBEBE"
50+
assert result["colors_dictionary"]["treat_none"] == "#BEBEBE"
51+
assert result["colors_dictionary"]["treat_all"] == "#BEBEBE"
52+
53+
54+
def test_time_curve_single_model_keeps_existing_black_style():
55+
curve_list = _curve_list(["model"], multiple_reference_groups=False)
56+
57+
result = _apply_color_values_times(curve_list, ["#123456"])
58+
59+
assert result["colors_dictionary"]["model"] == "#000000"
60+
61+
62+
def test_time_curve_custom_colors_require_one_per_reference_group():
63+
curve_list = _curve_list(["model_a", "model_b"])
64+
65+
with pytest.raises(ValueError, match="one color per reference group"):
66+
_apply_color_values_times(curve_list, ["#111111"])

0 commit comments

Comments
 (0)