diff --git a/src/rtichoke/_calibration_viz_spec_v2.py b/src/rtichoke/_calibration_viz_spec_v2.py index 9a5d4e15..f1224436 100644 --- a/src/rtichoke/_calibration_viz_spec_v2.py +++ b/src/rtichoke/_calibration_viz_spec_v2.py @@ -35,7 +35,9 @@ def _calibration_v2_spec_from_curve_list( f"Unsupported calibration type {calibration_type!r}; expected {supported}." ) - point_key = "deciles_dat" if calibration_type == "discrete" else "smooth_dat" + point_key = ( + "calibration_bins_dat" if calibration_type == "discrete" else "smooth_dat" + ) point_frame = _require_frame(calibration_curve_list, point_key) distribution_frame = _require_frame( calibration_curve_list, "histogram_for_calibration" diff --git a/src/rtichoke/_viz_browser.py b/src/rtichoke/_viz_browser.py index a5db2ee8..e0f711d5 100644 --- a/src/rtichoke/_viz_browser.py +++ b/src/rtichoke/_viz_browser.py @@ -67,7 +67,7 @@ def _calibration_spec_from_curve_list( calibration_curve_list: dict[str, Any], ) -> dict[str, object]: """Map existing discrete calibration output to the canonical spec.""" - rows = calibration_curve_list["deciles_dat"].select( + rows = calibration_curve_list["calibration_bins_dat"].select( "reference_group", "x", "y", "n_reals", "n" ) distribution_rows = calibration_curve_list["histogram_for_calibration"].select( diff --git a/src/rtichoke/calibration/__init__.py b/src/rtichoke/calibration/__init__.py index f792ad02..161c306d 100644 --- a/src/rtichoke/calibration/__init__.py +++ b/src/rtichoke/calibration/__init__.py @@ -2,6 +2,7 @@ Subpackage for Calibration """ +import functools import numpy as np from . import calibration as _calibration @@ -38,6 +39,7 @@ def _validate_outcome_values(reals, allowed_values): raise ValueError("Time-dependent outcomes must contain only 0, 1, and 2.") +@functools.wraps(_original_create_calibration_curve) def create_calibration_curve(*args, **kwargs): """Create an interactive calibration plot with a square main panel.""" probs = _argument(args, kwargs, "probs", 0) @@ -50,6 +52,7 @@ def create_calibration_curve(*args, **kwargs): ) +@functools.wraps(_original_create_calibration_curve_times) def create_calibration_curve_times(*args, **kwargs): """Create an interactive time-dependent calibration plot with a square main panel.""" probs = _argument(args, kwargs, "probs", 0) diff --git a/src/rtichoke/calibration/calibration.py b/src/rtichoke/calibration/calibration.py index 15879759..0baba7da 100644 --- a/src/rtichoke/calibration/calibration.py +++ b/src/rtichoke/calibration/calibration.py @@ -4,7 +4,6 @@ from typing import Any, Dict, List, Union, cast -# import pandas as pd import plotly.graph_objects as go from plotly.subplots import make_subplots from plotly.graph_objs._figure import Figure @@ -14,7 +13,16 @@ from smoothstate import smooth_state_lowess from ._secondary_cox import calculate_secondary_cox_smooth -# from rtichoke.helpers.send_post_request_to_r_rtichoke import send_requests_to_rtichoke_r + +def _validate_n_bins(n_bins: Any) -> int: + """Validates that n_bins is a positive integer >= 1.""" + if isinstance(n_bins, bool): + raise ValueError("n_bins must be a positive integer >= 1.") + if isinstance(n_bins, (int, np.integer)): + if n_bins < 1: + raise ValueError("n_bins must be a positive integer >= 1.") + return int(n_bins) + raise ValueError("n_bins must be a positive integer >= 1.") def create_calibration_curve( @@ -44,13 +52,15 @@ def create_calibration_curve( "#D1603D", "#585123", ], + *, + n_bins: int = 10, ) -> Figure: """Creates a Calibration Curve. This function generates a calibration curve, which evaluates how well the predicted probabilities from one or more models align with the observed binary outcomes. It can plot either discrete binned calibration - (deciles) or a smoothed calibration curve. + (10 bins by default) or a smoothed calibration curve. Parameters ---------- @@ -67,14 +77,22 @@ def create_calibration_curve( The width and height of the plot in pixels. Defaults to 600. color_values : List[str], optional A list of hex color strings for the plot lines/markers. + n_bins : int, optional + Number of bins for discrete calibration curves. Defaults to 10. Returns ------- Figure A Plotly ``Figure`` object representing the calibration curve. """ + n_bins = _validate_n_bins(n_bins) calibration_curve_list = _create_calibration_curve_list( - probs, reals, size=size, color_values=color_values + probs, + reals, + calibration_type=calibration_type, + size=size, + color_values=color_values, + n_bins=n_bins, ) calibration_curve = _create_plotly_curve_from_calibration_curve_list( @@ -116,6 +134,8 @@ def create_calibration_curve_times( "#D1603D", "#585123", ], + *, + n_bins: int = 10, ) -> Figure: """Create a time-dependent calibration curve across fixed horizons. @@ -151,6 +171,8 @@ def create_calibration_curve_times( Width and height of the Plotly figure in pixels. Defaults to 600. color_values : List[str], optional List of hex color strings for traces. + n_bins : int, optional + Number of bins for discrete calibration curves. Defaults to 10. Returns ------- @@ -162,6 +184,7 @@ def create_calibration_curve_times( ValueError If a heuristic set requests `competing_heuristic='adjusted_as_censored'`. """ + n_bins = _validate_n_bins(n_bins) unsupported_competing_as_censored = any( heuristics.get("competing_heuristic") == "adjusted_as_censored" @@ -186,6 +209,7 @@ def create_calibration_curve_times( bandwidth=bandwidth, size=size, color_values=color_values, + n_bins=n_bins, ) fig = _create_plotly_curve_from_calibration_curve_list_times( @@ -233,7 +257,7 @@ def _create_plotly_curve_from_calibration_curve_list_times( # Calibration curve (discrete or smooth) if calibration_type == "discrete": - data_subset = calibration_curve_list["deciles_dat"].filter( + data_subset = calibration_curve_list["calibration_bins_dat"].filter( (pl.col("reference_group") == group) & (pl.col("fixed_time_horizon") == horizon) ) @@ -344,15 +368,7 @@ def _create_plotly_curve_from_calibration_curve_list_times( def _create_plotly_curve_from_calibration_curve_list( calibration_curve_list: Dict[str, Any], calibration_type: str = "discrete" ) -> Figure: - """Create plotly curve from calibration curve list - - Args: - calibration_curve_list (Dict[str, Any]): _description_ - calibration_type (str, optional): _description_. Defaults to "discrete". - - Returns: - Figure: _description_ - """ + """Create plotly curve from calibration curve list""" calibration_curve = make_subplots( rows=2, cols=1, shared_xaxes=True, x_title="Predicted", row_heights=[0.8, 0.2] ) @@ -404,15 +420,15 @@ def _create_plotly_curve_from_calibration_curve_list( if k != "reference_line" ] for reference_group in reference_groups: - dec_sub = calibration_curve_list["deciles_dat"].filter( + bin_sub = calibration_curve_list["calibration_bins_dat"].filter( pl.col("reference_group") == reference_group ) calibration_curve.add_trace( go.Scatter( - x=dec_sub.get_column("x").to_list(), - y=dec_sub.get_column("y").to_list(), - hovertext=dec_sub.get_column("text").to_list(), + x=bin_sub.get_column("x").to_list(), + y=bin_sub.get_column("y").to_list(), + hovertext=bin_sub.get_column("text").to_list(), name=reference_group, legendgroup=reference_group, hoverinfo="text", @@ -541,11 +557,13 @@ def _create_plotly_curve_from_calibration_curve_list( return calibration_curve -def _make_deciles_dat_binary( +def _make_calibration_bins_dat_binary( probs: Dict[str, np.ndarray], reals: Union[np.ndarray, Dict[str, np.ndarray]], n_bins: int = 10, ) -> pl.DataFrame: + n_bins = _validate_n_bins(n_bins) + if isinstance(reals, dict): frames: list[pl.DataFrame] = [] @@ -630,21 +648,56 @@ def _make_deciles_dat_binary( df = pl.concat(frames, how="vertical") - df = df.with_columns( - [ - pl.col("prob").cast(pl.Float64), - pl.col("real").cast(pl.Float64), - ( - (pl.col("prob").rank("ordinal").over(["reference_group", "model"]) - 1) - * n_bins - // pl.len().over(["reference_group", "model"]) - + 1 - ).alias("decile"), - ] - ) + # Apply dplyr::ntile() bucket allocation per (reference_group, model) + group_cols = ["reference_group", "model"] + partitioned_frames = [] + + for key, sub_df in df.group_by(group_cols, maintain_order=True): + reference_group, model = key + N = sub_df.height + p_vals = sub_df["prob"].to_numpy() + + if N == 0: + sub_df_with_bin = sub_df.with_columns( + pl.lit(1, dtype=pl.Int64).alias("bin") + ) + elif len(np.unique(p_vals)) == 1: + # All predictions identical -> one aggregate calibration bin = 1 + sub_df_with_bin = sub_df.with_columns( + pl.lit(1, dtype=pl.Int64).alias("bin") + ) + elif n_bins > N: + # B > N -> occupied bin labels are 1..N, one observation each in stable order + ord_ranks = sub_df["prob"].rank("ordinal").to_numpy().astype(int) + sub_df_with_bin = sub_df.with_columns( + pl.Series("bin", ord_ranks, dtype=pl.Int64) + ) + else: + # B <= N -> dplyr::ntile() semantics: + # stable ordinal rank 1..N + ord_ranks = sub_df["prob"].rank("ordinal").to_numpy().astype(int) + q = N // n_bins + rem = N % n_bins + # first rem bins have size q + 1, remaining B - rem bins have size q + # build lookup array bin_for_rank mapping rank (1-indexed) -> bin (1-indexed) + bin_for_rank = np.empty(N + 1, dtype=int) + curr_rank = 1 + for b in range(1, n_bins + 1): + bin_size = q + 1 if b <= rem else q + bin_for_rank[curr_rank : curr_rank + bin_size] = b + curr_rank += bin_size + + assigned_bins = bin_for_rank[ord_ranks] + sub_df_with_bin = sub_df.with_columns( + pl.Series("bin", assigned_bins, dtype=pl.Int64) + ) + + partitioned_frames.append(sub_df_with_bin) - deciles_data = ( - df.group_by(["reference_group", "model", "decile"]) + df = pl.concat(partitioned_frames) + + calibration_bins_data = ( + df.group_by(["reference_group", "model", "bin"]) .agg( [ pl.len().alias("n"), @@ -653,10 +706,10 @@ def _make_deciles_dat_binary( pl.sum("real").alias("n_reals"), ] ) - .sort(["reference_group", "model", "decile"]) + .sort(["reference_group", "model", "bin"]) ) - return deciles_data + return calibration_bins_data def _check_performance_type_by_probs_and_reals( @@ -695,13 +748,20 @@ def _create_calibration_curve_list( "#D1603D", "#585123", ], + *, + calibration_type: str = "discrete", + n_bins: int = 10, ) -> Dict[str, Any]: - deciles_data = _make_deciles_dat_binary(probs, reals) + n_bins = _validate_n_bins(n_bins) + effective_n_bins = n_bins if calibration_type == "discrete" else 10 + calibration_bins_data = _make_calibration_bins_dat_binary( + probs, reals, n_bins=effective_n_bins + ) performance_type = _check_performance_type_by_probs_and_reals(probs, reals) smooth_dat = _calculate_smooth_curve(probs, reals, performance_type) - deciles_data, smooth_dat = _add_hover_text_to_calibration_data( - deciles_data, smooth_dat, performance_type + calibration_bins_data, smooth_dat = _add_hover_text_to_calibration_data( + calibration_bins_data, smooth_dat, performance_type ) reference_data = _create_reference_data_for_calibration_curve() @@ -714,17 +774,14 @@ def _create_calibration_curve_list( histogram_for_calibration = _create_histogram_for_calibration(probs) - limits = _define_limits_for_calibration_plot(deciles_data) + limits = _define_limits_for_calibration_plot(calibration_bins_data) axes_ranges = {"xaxis": limits, "yaxis": limits} - smooth_dat = _calculate_smooth_curve(probs, reals, performance_type) - calibration_curve_list = { - "deciles_dat": deciles_data, + "calibration_bins_dat": calibration_bins_data, "smooth_dat": smooth_dat, "reference_data": reference_data, "histogram_for_calibration": histogram_for_calibration, - # "histogram_opacity": [0.4], "axes_ranges": axes_ranges, "colors_dictionary": colors_dictionary, "performance_type": [performance_type], @@ -776,8 +833,6 @@ def process_single_array(p, r, group_name): if isinstance(reals, dict): for model_name, prob_array in probs.items(): - # This logic assumes that for multiple populations, one model's probs are evaluated against multiple real outcomes. - # This might need adjustment based on the exact structure for multiple models and populations. if len(probs) == 1 and len(reals) > 1: # One model, multiple populations for pop_name, real_array in reals.items(): frame = process_single_array(prob_array, real_array, pop_name) @@ -845,13 +900,13 @@ def process_single_array(p, r, group_name): def _add_hover_text_to_calibration_data( - deciles_dat: pl.DataFrame, + calibration_bins_dat: pl.DataFrame, smooth_dat: pl.DataFrame, performance_type: str, ) -> tuple[pl.DataFrame, pl.DataFrame]: - """Adds hover text to the deciles and smooth dataframes.""" + """Adds hover text to the calibration bins and smooth dataframes.""" if performance_type != "one model": - deciles_dat = deciles_dat.with_columns( + calibration_bins_dat = calibration_bins_dat.with_columns( pl.concat_str( [ pl.lit(""), @@ -881,7 +936,7 @@ def _add_hover_text_to_calibration_data( ).alias("text") ) else: - deciles_dat = deciles_dat.with_columns( + calibration_bins_dat = calibration_bins_dat.with_columns( pl.concat_str( [ pl.lit("Predicted: "), @@ -906,7 +961,7 @@ def _add_hover_text_to_calibration_data( ] ).alias("text") ) - return deciles_dat, smooth_dat + return calibration_bins_dat, smooth_dat def _create_colors_dictionary_for_calibration( @@ -951,23 +1006,25 @@ def _create_histogram_for_calibration(probs: Dict[str, np.ndarray]) -> pl.DataFr return histogram_for_calibration -def _define_limits_for_calibration_plot(deciles_dat: pl.DataFrame) -> List[float]: - if deciles_dat.height == 1: +def _define_limits_for_calibration_plot( + calibration_bins_dat: pl.DataFrame, +) -> List[float]: + if calibration_bins_dat.height == 1: lower_bound, upper_bound = 0.0, 1.0 else: lower_bound = float( max( 0, min( - cast(float, deciles_dat["x"].min()), - cast(float, deciles_dat["y"].min()), + cast(float, calibration_bins_dat["x"].min()), + cast(float, calibration_bins_dat["y"].min()), ), ) ) upper_bound = float( max( - cast(float, deciles_dat["x"].max()), - cast(float, deciles_dat["y"].max()), + cast(float, calibration_bins_dat["x"].max()), + cast(float, calibration_bins_dat["y"].max()), ) ) @@ -1138,35 +1195,38 @@ def _aj_risk_at_horizon(df: pl.DataFrame, horizon: float) -> float: return float(estimate["state_occupancy_probability_1"][0]) -def _make_adjusted_deciles_data( +def _make_adjusted_calibration_bins_data( df: pl.DataFrame, horizon: float, n_bins: int = 10 ) -> pl.DataFrame: """Create calibration groups using within-group Aalen-Johansen risks.""" + n_bins = _validate_n_bins(n_bins) grouped = df.with_columns( ( (pl.col("prob").rank("average").over("reference_group") - 1) * n_bins // pl.len().over("reference_group") + 1 - ).alias("decile") + ) + .cast(pl.Int64) + .alias("bin") ) rows = [] - for key, group_df in grouped.group_by(["reference_group", "decile"]): - reference_group, decile = key + for key, group_df in grouped.group_by(["reference_group", "bin"]): + reference_group, bin_val = key estimate = _aj_risk_at_horizon(group_df, horizon) n = group_df.height rows.append( { "reference_group": reference_group, "model": reference_group, - "decile": decile, + "bin": int(bin_val), "n": n, "x": cast(float, group_df["prob"].mean()), "y": estimate, "n_reals": estimate * n, } ) - return pl.DataFrame(rows).sort(["reference_group", "decile"]) + return pl.DataFrame(rows).sort(["reference_group", "bin"]) def _calculate_local_aj_smooth( @@ -1291,14 +1351,19 @@ def _create_calibration_curve_list_times( "#D1603D", "#585123", ], + *, + n_bins: int = 10, ) -> Dict[str, Any]: """ Creates the data structures needed for a time-dependent calibration curve plot. """ + n_bins = _validate_n_bins(n_bins) + effective_n_bins = n_bins if calibration_type == "discrete" else 10 + # Part 1: Prepare initial dataframe from inputs initial_df = _build_initial_df_for_times(probs, reals, times) # Part 2: Iterate and generate calibration data for each horizon/heuristic - all_deciles = [] + all_calibration_bins = [] all_smooth = [] all_histograms = [] @@ -1319,7 +1384,9 @@ def _create_calibration_curve_list_times( if df_adj.height == 0: continue - deciles_data = _make_adjusted_deciles_data(df_adj, horizon) + calibration_bins_data = _make_adjusted_calibration_bins_data( + df_adj, horizon, n_bins=effective_n_bins + ) probs_adj = { key[0]: group_df["prob"].to_numpy() for key, group_df in df_adj.group_by( @@ -1351,11 +1418,13 @@ def _create_calibration_curve_list_times( "Supported options are 'local_aj', 'secondary_cox', and 'pseudo_values'." ) else: - smooth_data = deciles_data.select("x", "y", "reference_group") + smooth_data = calibration_bins_data.select( + "x", "y", "reference_group" + ) hist_data = _create_histogram_for_calibration(probs_adj) - all_deciles.append( - deciles_data.with_columns( + all_calibration_bins.append( + calibration_bins_data.with_columns( pl.lit(horizon).alias("fixed_time_horizon") ) ) @@ -1389,10 +1458,14 @@ def _create_calibration_curve_list_times( if not isinstance(reals, dict) and len(probs) == 1: reals_adj = next(iter(reals_adj.values())) - # Deciles - deciles_data = _make_deciles_dat_binary(probs_adj, reals_adj) - all_deciles.append( - deciles_data.with_columns(pl.lit(horizon).alias("fixed_time_horizon")) + # Calibration bins + calibration_bins_data = _make_calibration_bins_dat_binary( + probs_adj, reals_adj, n_bins=effective_n_bins + ) + all_calibration_bins.append( + calibration_bins_data.with_columns( + pl.lit(horizon).alias("fixed_time_horizon") + ) ) # Smooth curve @@ -1418,7 +1491,7 @@ def _create_calibration_curve_list_times( "Supported options are 'local_aj', 'secondary_cox', and 'pseudo_values'." ) else: - smooth_data = deciles_data.select("x", "y", "reference_group") + smooth_data = calibration_bins_data.select("x", "y", "reference_group") all_smooth.append( smooth_data.with_columns(pl.lit(horizon).alias("fixed_time_horizon")) ) @@ -1430,17 +1503,17 @@ def _create_calibration_curve_list_times( ) # Part 3: Combine results and create final dictionary - if not all_deciles: + if not all_calibration_bins: raise ValueError( "No data remaining after applying heuristics and time horizons." ) - deciles_dat_final = pl.concat(all_deciles) + calibration_bins_dat_final = pl.concat(all_calibration_bins) smooth_dat_final = pl.concat(all_smooth) histogram_final = pl.concat(all_histograms) # Add hover text - deciles_dat_final, smooth_dat_final = _add_hover_text_to_calibration_data( - deciles_dat_final, smooth_dat_final, performance_type + calibration_bins_dat_final, smooth_dat_final = _add_hover_text_to_calibration_data( + calibration_bins_dat_final, smooth_dat_final, performance_type ) reference_data = _create_reference_data_for_calibration_curve() @@ -1448,11 +1521,11 @@ def _create_calibration_curve_list_times( colors_dictionary = _create_colors_dictionary_for_calibration( reference_groups, color_values, performance_type ) - limits = _define_limits_for_calibration_plot(deciles_dat_final) + limits = _define_limits_for_calibration_plot(calibration_bins_dat_final) axes_ranges = {"xaxis": limits, "yaxis": limits} calibration_curve_list = { - "deciles_dat": deciles_dat_final, + "calibration_bins_dat": calibration_bins_dat_final, "smooth_dat": smooth_dat_final, "reference_data": reference_data, "histogram_for_calibration": histogram_final, diff --git a/tests/test_calibration_bins.py b/tests/test_calibration_bins.py new file mode 100644 index 00000000..525fbab9 --- /dev/null +++ b/tests/test_calibration_bins.py @@ -0,0 +1,338 @@ +"""Tests for calibration bins (n_bins) cross-language parity, validation, and contracts.""" + +import inspect +import numpy as np +import polars as pl +import pytest +from plotly.graph_objs._figure import Figure + +import rtichoke +from rtichoke import ( + create_calibration_curve, + create_calibration_curve_times, +) +from rtichoke.calibration.calibration import ( + _create_calibration_curve_list, + _create_calibration_curve_list_times, + _make_calibration_bins_dat_binary, + _make_adjusted_calibration_bins_data, +) +from rtichoke._calibration_viz_spec_v2 import _calibration_v2_spec_from_curve_list +from rtichoke._viz_browser import _calibration_spec_from_curve_list +from rtichoke.processing.evaluation_semantics import _EvaluationMetadata + + +def test_top_level_export_signatures_expose_n_bins(): + sig = inspect.signature(rtichoke.create_calibration_curve) + assert "n_bins" in sig.parameters + param = sig.parameters["n_bins"] + assert param.kind == inspect.Parameter.KEYWORD_ONLY + assert param.default == 10 + + sig_times = inspect.signature(rtichoke.create_calibration_curve_times) + assert "n_bins" in sig_times.parameters + param_times = sig_times.parameters["n_bins"] + assert param_times.kind == inspect.Parameter.KEYWORD_ONLY + assert param_times.default == 10 + + +def test_create_calibration_curve_list_positional_compatibility(): + p = np.linspace(0.1, 0.9, 10) + y = np.tile([0, 1], 5) + cl = _create_calibration_curve_list({"m1": p}, y, 800) + assert cl["size"] == [(800, 800)] + + +def test_n_bins_validation_invalid_values(): + probs = {"m1": np.array([0.1, 0.4, 0.7, 0.9])} + reals = np.array([0, 0, 1, 1]) + + invalid_values = [0, -1, -10, 2.5, None, float("nan"), "10", True, False] + for val in invalid_values: + with pytest.raises(ValueError, match="n_bins must be a positive integer >= 1."): + create_calibration_curve(probs, reals, n_bins=val) + + with pytest.raises(ValueError, match="n_bins must be a positive integer >= 1."): + create_calibration_curve_times( + probs, + reals, + times=np.array([1, 2, 3, 4]), + fixed_time_horizons=[2.0], + heuristics_sets=[ + { + "censoring_heuristic": "adjusted", + "competing_heuristic": "adjusted_as_negative", + } + ], + n_bins=val, + ) + + +def test_n_bins_validation_accepts_numpy_integers(): + probs = {"m1": np.array([0.1, 0.4, 0.7, 0.9])} + reals = np.array([0, 0, 1, 1]) + fig = create_calibration_curve(probs, reals, n_bins=np.int64(5)) + assert isinstance(fig, Figure) + + +def test_r_parity_n12_b10(): + # N=12, B=10 -> sizes must be 2, 2, 1, 1, 1, 1, 1, 1, 1, 1 + p = np.linspace(0.01, 0.99, 12) + y = np.array([0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1]) + probs = {"m1": p} + reals = y + + bins_dat = _make_calibration_bins_dat_binary(probs, reals, n_bins=10) + assert bins_dat.height == 10 + counts = bins_dat["n"].to_list() + assert counts == [2, 2, 1, 1, 1, 1, 1, 1, 1, 1] + assert bins_dat["bin"].to_list() == list(range(1, 11)) + + +def test_r_parity_n11_b10(): + # N=11, B=10 -> sizes must be 2, 1, 1, 1, 1, 1, 1, 1, 1, 1 + p = np.linspace(0.01, 0.99, 11) + y = np.array([0, 1] * 5 + [0]) + probs = {"m1": p} + reals = y + + bins_dat = _make_calibration_bins_dat_binary(probs, reals, n_bins=10) + assert bins_dat.height == 10 + counts = bins_dat["n"].to_list() + assert counts == [2, 1, 1, 1, 1, 1, 1, 1, 1, 1] + + +def test_r_parity_n5_b10(): + # N=5, B=10 -> occupied labels 1, 2, 3, 4, 5, one observation each + p = np.linspace(0.1, 0.9, 5) + y = np.array([0, 0, 1, 0, 1]) + probs = {"m1": p} + reals = y + + bins_dat = _make_calibration_bins_dat_binary(probs, reals, n_bins=10) + assert bins_dat.height == 5 + assert bins_dat["bin"].to_list() == [1, 2, 3, 4, 5] + assert bins_dat["n"].to_list() == [1, 1, 1, 1, 1] + + +def test_r_parity_b_greater_than_n(): + p = np.array([0.1, 0.3, 0.8]) + y = np.array([0, 1, 1]) + bins_dat = _make_calibration_bins_dat_binary({"m1": p}, y, n_bins=5) + assert bins_dat.height == 3 + assert bins_dat["bin"].to_list() == [1, 2, 3] + assert bins_dat["n"].to_list() == [1, 1, 1] + + +def test_r_parity_all_predictions_identical(): + p = np.array([0.5, 0.5, 0.5, 0.5, 0.5]) + y = np.array([0, 1, 0, 1, 0]) + bins_dat = _make_calibration_bins_dat_binary({"m1": p}, y, n_bins=10) + assert bins_dat.height == 1 + assert bins_dat["bin"].to_list() == [1] + assert bins_dat["n"].to_list() == [5] + assert bins_dat["x"].to_list() == [0.5] + assert bins_dat["y"].to_list() == [0.4] + + +def test_r_parity_partial_ties_stable_ordinal_order(): + p = np.array([0.1, 0.1, 0.2, 0.3, 0.4, 0.5]) + y = np.array([0, 1, 0, 1, 0, 1]) + # N=6, B=3 -> q=2, rem=0 -> 3 bins of size 2 + bins_dat = _make_calibration_bins_dat_binary({"m1": p}, y, n_bins=3) + assert bins_dat.height == 3 + assert bins_dat["n"].to_list() == [2, 2, 2] + assert bins_dat["bin"].to_list() == [1, 2, 3] + + +def test_static_default_and_explicit_n_bins(): + p = np.linspace(0.01, 0.99, 100) + y = (p > 0.5).astype(int) + probs = {"m1": p} + + # Default n_bins = 10 + cl_default = _create_calibration_curve_list(probs, y) + assert cl_default["calibration_bins_dat"].height == 10 + + # Explicit n_bins = 8 + cl_8 = _create_calibration_curve_list(probs, y, n_bins=8) + assert cl_8["calibration_bins_dat"].height == 8 + + # n_bins = 1 + cl_1 = _create_calibration_curve_list(probs, y, n_bins=1) + assert cl_1["calibration_bins_dat"].height == 1 + assert cl_1["calibration_bins_dat"]["n"].to_list() == [100] + + +def test_multiple_models_and_populations_preserve_identities(): + p1 = np.linspace(0.1, 0.9, 20) + p2 = np.linspace(0.2, 0.8, 20) + y = np.tile([0, 1], 10) + + cl_multi = _create_calibration_curve_list( + {"Model A": p1, "Model B": p2}, y, n_bins=5 + ) + df_bins = cl_multi["calibration_bins_dat"] + assert set(df_bins["reference_group"].unique().to_list()) == {"Model A", "Model B"} + assert set(df_bins["model"].unique().to_list()) == {"Model A", "Model B"} + assert df_bins.filter(pl.col("reference_group") == "Model A").height == 5 + assert df_bins.filter(pl.col("reference_group") == "Model B").height == 5 + + +def test_non_adjusted_time_inherits_static_bin_contract(): + probs = {"m1": np.linspace(0.01, 0.99, 12)} + reals = np.tile([0, 1], 6) + times = np.full(12, 3.0) # follow-up time 3.0 > horizon 2.0 -> no censoring + + cl_time = _create_calibration_curve_list_times( + probs, + reals, + times, + fixed_time_horizons=[2.0], + heuristics_sets=[ + {"censoring_heuristic": "excluded", "competing_heuristic": "excluded"} + ], + n_bins=10, + ) + df_bins = cl_time["calibration_bins_dat"] + assert df_bins.height == 10 + assert df_bins["n"].to_list() == [2, 2, 1, 1, 1, 1, 1, 1, 1, 1] + + +def test_adjusted_time_calibration_preserved_semantics_and_ties(): + # Test fixture with partially tied predictions: + # 3 identical predictions at 0.2, 3 identical at 0.5, 4 distinct at 0.6, 0.7, 0.8, 0.9 + p_tied = [0.2, 0.2, 0.2, 0.5, 0.5, 0.5, 0.6, 0.7, 0.8, 0.9] + y_vals = [0, 1, 0, 1, 0, 1, 0, 1, 0, 1] + times = [1.0] * 10 + df = pl.DataFrame( + { + "reference_group": ["m1"] * 10, + "prob": p_tied, + "real": y_vals, + "time": times, + } + ) + + # Compute with n_bins=10 + res_10 = _make_adjusted_calibration_bins_data(df, horizon=2.0, n_bins=10) + assert "bin" in res_10.columns + assert "decile" not in res_10.columns + assert res_10["bin"].dtype in (pl.Int64, pl.Int32) + + # All tied observations (first 3 at 0.2) must belong to the exact same bin + grouped_10 = df.with_columns( + ( + (pl.col("prob").rank("average").over("reference_group") - 1) + * 10 + // pl.len().over("reference_group") + + 1 + ) + .cast(pl.Int64) + .alias("bin") + ) + first_3_bins = grouped_10["bin"][:3].to_list() + assert len(set(first_3_bins)) == 1, ( + "Tied predictions 0.2 must share the exact same bin" + ) + + mid_3_bins = grouped_10["bin"][3:6].to_list() + assert len(set(mid_3_bins)) == 1, ( + "Tied predictions 0.5 must share the exact same bin" + ) + + # Compute with n_bins=5 + res_5 = _make_adjusted_calibration_bins_data(df, horizon=2.0, n_bins=5) + assert "bin" in res_5.columns + assert res_5["bin"].dtype in (pl.Int64, pl.Int32) + + grouped_5 = df.with_columns( + ( + (pl.col("prob").rank("average").over("reference_group") - 1) + * 5 + // pl.len().over("reference_group") + + 1 + ) + .cast(pl.Int64) + .alias("bin") + ) + first_3_bins_5 = grouped_5["bin"][:3].to_list() + assert len(set(first_3_bins_5)) == 1, ( + "Tied predictions 0.2 must share the exact same bin under n_bins=5" + ) + + +def test_smooth_static_n_bins_no_effect(): + p = np.linspace(0.01, 0.99, 50) + y = np.tile([0, 1], 25) + probs = {"m1": p} + + cl_smooth_10 = _create_calibration_curve_list( + probs, y, calibration_type="smooth", n_bins=10 + ) + cl_smooth_5 = _create_calibration_curve_list( + probs, y, calibration_type="smooth", n_bins=5 + ) + + assert cl_smooth_10["smooth_dat"].equals(cl_smooth_5["smooth_dat"]) + assert cl_smooth_10["axes_ranges"] == cl_smooth_5["axes_ranges"] + + +def test_smooth_time_n_bins_no_effect(): + probs = {"m1": np.linspace(0.01, 0.99, 50)} + reals = np.tile([0, 1], 25) + times = np.ones(50) + + cl_time_10 = _create_calibration_curve_list_times( + probs, + reals, + times, + fixed_time_horizons=[2.0], + heuristics_sets=[ + { + "censoring_heuristic": "adjusted", + "competing_heuristic": "adjusted_as_negative", + } + ], + calibration_type="smooth", + n_bins=10, + ) + cl_time_5 = _create_calibration_curve_list_times( + probs, + reals, + times, + fixed_time_horizons=[2.0], + heuristics_sets=[ + { + "censoring_heuristic": "adjusted", + "competing_heuristic": "adjusted_as_negative", + } + ], + calibration_type="smooth", + n_bins=5, + ) + + assert cl_time_10["smooth_dat"].equals(cl_time_5["smooth_dat"]) + assert cl_time_10["axes_ranges"] == cl_time_5["axes_ranges"] + + +def test_v1_and_v2_canonical_adapters_consume_calibration_bins_dat(): + p = np.linspace(0.1, 0.9, 10) + y = np.tile([0, 1], 5) + probs = {"m1": p} + reals = y + + curve_list = _create_calibration_curve_list(probs, reals, n_bins=5) + assert "calibration_bins_dat" in curve_list + assert "deciles_dat" not in curve_list + + # v1 adapter + v1_spec = _calibration_spec_from_curve_list(curve_list) + assert v1_spec["type"] == "calibration" + assert len(v1_spec["data"]) == 5 + + # v2 adapter + meta = {"m1": _EvaluationMetadata("m1", "eval1", "m1", "Pop 1")} + v2_spec = _calibration_v2_spec_from_curve_list(curve_list, meta) + assert v2_spec["type"] == "calibration" + assert len(v2_spec["data"]) == 5 diff --git a/tests/test_calibration_v2.py b/tests/test_calibration_v2.py index 25d3b34e..22d8d094 100644 --- a/tests/test_calibration_v2.py +++ b/tests/test_calibration_v2.py @@ -85,7 +85,7 @@ def test_calibration_v2_discrete_values_pass_through_without_recalculation(): ) production_rows = ( - curve_list["deciles_dat"] + curve_list["calibration_bins_dat"] .select("reference_group", "x", "y", "n_reals", "n") .to_dicts() ) diff --git a/tests/test_semantic_characterization.py b/tests/test_semantic_characterization.py index 84f65564..562ef575 100644 --- a/tests/test_semantic_characterization.py +++ b/tests/test_semantic_characterization.py @@ -197,7 +197,9 @@ def test_calibration_keeps_grouping_and_one_global_identity_line(): reals=REALS_EQUAL, ) assert multi_model["performance_type"] == ["multiple models"] - assert set(multi_model["deciles_dat"]["reference_group"].unique().to_list()) == { + assert set( + multi_model["calibration_bins_dat"]["reference_group"].unique().to_list() + ) == { "Model A", "Model B", } @@ -209,7 +211,7 @@ def test_calibration_keeps_grouping_and_one_global_identity_line(): ) assert multi_population["performance_type"] == ["multiple populations"] assert set( - multi_population["deciles_dat"]["reference_group"].unique().to_list() + multi_population["calibration_bins_dat"]["reference_group"].unique().to_list() ) == {"Population low", "Population high"} assert multi_population["reference_data"].height == 101 diff --git a/user_guide/05-calibration-curves.qmd b/user_guide/05-calibration-curves.qmd index 72fdcd99..78ceaba5 100644 --- a/user_guide/05-calibration-curves.qmd +++ b/user_guide/05-calibration-curves.qmd @@ -19,7 +19,7 @@ fig = rk.create_calibration_curve( ) ``` -`calibration_type` can be set to `"discrete"` (deciles) or `"smooth"` (lowess curve). +`calibration_type` can be set to `"discrete"` (binned calibration, 10 bins by default) or `"smooth"` (lowess curve). ## Time-Dependent Calibration (`create_calibration_curve_times`) @@ -72,7 +72,18 @@ When `calibration_type="smooth"`, `create_calibration_curve_times` supports thre ## Discrete (Binned) Calibration -For binned decile plots at time horizons, pass `calibration_type="discrete"`: +For binned calibration plots, pass `calibration_type="discrete"`. The number of bins defaults to 10 (`n_bins=10`), but can be customized using the `n_bins` parameter (e.g., `n_bins=8`): + +```python +fig = rk.create_calibration_curve( + probs={"Model A": probs_a}, + reals=reals_binary, + calibration_type="discrete", + n_bins=8, +) +``` + +Similarly, for time-dependent binned calibration: ```python fig = rk.create_calibration_curve_times( @@ -82,5 +93,6 @@ fig = rk.create_calibration_curve_times( fixed_time_horizons=[5.0], heuristics_sets=heuristics_sets, calibration_type="discrete", + n_bins=8, ) ```