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
4 changes: 2 additions & 2 deletions .github/workflows/python-package.yml
Original file line number Diff line number Diff line change
Expand Up @@ -52,14 +52,14 @@ jobs:
prefix = "rtichoke/_vendor/rtichoke_viz/"
required = {
f"{prefix}VENDORED_FROM",
f"{prefix}rtichoke-viz-0.5.0.tar.gz",
f"{prefix}rtichoke-viz-0.6.0.tar.gz",
f"{prefix}rtichoke-viz.js",
f"{prefix}rtichoke-viz.css",
f"{prefix}rtichoke-viz.schema.json",
f"{prefix}rtichoke-viz-v2.schema.json",
}
assert required <= names
assert f"{prefix}rtichoke-viz-0.4.0.tar.gz" not in names
assert f"{prefix}rtichoke-viz-0.5.0.tar.gz" not in names
PY

- name: Run tests
Expand Down
15 changes: 13 additions & 2 deletions .github/workflows/quarto-acceptance.yml
Original file line number Diff line number Diff line change
Expand Up @@ -7,19 +7,27 @@ on:
- "src/rtichoke/summary_report/**"
- "src/rtichoke/_report_browser.py"
- "src/rtichoke/_report_spec.py"
- "src/rtichoke/_renderers.py"
- "src/rtichoke/_decision_curve_viz_spec_v2.py"
- "src/rtichoke/utility/decision.py"
- "src/rtichoke/_vendor/rtichoke_viz/**"
- "tests/test_quarto_summary_report_browser.py"
- "tests/test_summary_report_browser.py"
- "tests/test_decision_curve_browser_acceptance.py"
- ".github/workflows/quarto-acceptance.yml"
pull_request:
branches: ["main"]
paths:
- "src/rtichoke/summary_report/**"
- "src/rtichoke/_report_browser.py"
- "src/rtichoke/_report_spec.py"
- "src/rtichoke/_renderers.py"
- "src/rtichoke/_decision_curve_viz_spec_v2.py"
- "src/rtichoke/utility/decision.py"
- "src/rtichoke/_vendor/rtichoke_viz/**"
- "tests/test_quarto_summary_report_browser.py"
- "tests/test_summary_report_browser.py"
- "tests/test_decision_curve_browser_acceptance.py"
- ".github/workflows/quarto-acceptance.yml"
workflow_dispatch:

Expand Down Expand Up @@ -48,5 +56,8 @@ jobs:
- name: Install Playwright Chromium for browser acceptance
run: uv run --with playwright playwright install chromium

- name: Run Quarto browser report acceptance tests
run: uv run --with playwright pytest tests/test_quarto_summary_report_browser.py
- name: Run browser acceptance tests
run: >-
uv run --with playwright pytest
tests/test_quarto_summary_report_browser.py
tests/test_decision_curve_browser_acceptance.py
174 changes: 174 additions & 0 deletions src/rtichoke/_decision_curve_viz_spec_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
"""Canonical static Decision Curve v2 adapter.

This module translates already-computed production Decision Curve quantities
into the shared rtichoke_viz contract. It deliberately does not recompute model
statistics or threshold membership.
"""

from __future__ import annotations

from collections.abc import Mapping

import polars as pl

from rtichoke.processing.evaluation_semantics import _EvaluationMetadata

_REQUIRED_COLUMNS = {
"reference_group",
"chosen_cutoff",
"net_benefit",
"real_positives",
"n",
}


def _decision_curve_v2_spec_from_performance_data(
performance_data: pl.DataFrame,
evaluation_metadata: Mapping[str, _EvaluationMetadata],
*,
min_p_threshold: float = 0.0,
max_p_threshold: float = 1.0,
) -> dict[str, object]:
"""Build canonical static Decision Curve v2 from production quantities."""
missing = _REQUIRED_COLUMNS.difference(performance_data.columns)
if missing:
raise ValueError(
"Decision Curve performance data is missing columns: "
+ ", ".join(sorted(missing))
)

rows = (
performance_data.filter(
pl.col("chosen_cutoff").is_finite() & pl.col("net_benefit").is_finite()
)
.select(
"reference_group",
"chosen_cutoff",
"net_benefit",
"real_positives",
"n",
)
.to_dicts()
)

row_groups = {str(row["reference_group"]) for row in rows}
missing_metadata = row_groups.difference(evaluation_metadata)
if missing_metadata:
raise ValueError(
"Decision Curve rows are missing evaluation metadata: "
+ ", ".join(sorted(missing_metadata))
)

ordered_groups = [group for group in evaluation_metadata if group in row_groups]
evaluation_ids = {
group: f"evaluation-{index}"
for index, group in enumerate(ordered_groups, start=1)
}
series_ids = {
group: f"series-{index}" for index, group in enumerate(ordered_groups, start=1)
}

evaluations: list[dict[str, object]] = []
series: list[dict[str, object]] = []
for group in ordered_groups:
metadata = evaluation_metadata[group]
evaluation: dict[str, object] = {
"id": evaluation_ids[group],
"population": metadata.population,
}
if metadata.model is not None:
evaluation["model"] = metadata.model
evaluations.append(evaluation)

display_value = metadata.model or metadata.population
series.append(
{
"id": series_ids[group],
"evaluationId": evaluation_ids[group],
"display": {
"label": display_value,
"group": display_value,
"role": "model" if metadata.model is not None else "population",
},
}
)

data = [
{
"seriesId": series_ids[str(row["reference_group"])],
"threshold": float(row["chosen_cutoff"]),
"netBenefit": float(row["net_benefit"]),
}
for row in rows
]

prevalence_values: dict[str, set[float]] = {}
population_thresholds: dict[str, list[float]] = {}
for row in rows:
group = str(row["reference_group"])
population = evaluation_metadata[group].population
n = float(row["n"])
if n <= 0:
raise ValueError("Decision Curve population size must be positive.")
prevalence_values.setdefault(population, set()).add(
float(row["real_positives"]) / n
)
threshold = float(row["chosen_cutoff"])
if 0.0 <= threshold < 1.0:
population_thresholds.setdefault(population, []).append(threshold)

populations = list(
dict.fromkeys(metadata.population for metadata in evaluation_metadata.values())
)
references: list[dict[str, object]] = [
{
"type": "horizontal",
"scope": "global",
"value": 0.0,
"label": "Treat None",
"benchmark": "treat_none",
}
]
for population in populations:
values = prevalence_values.get(population, set())
if not values:
continue
if len(values) != 1:
raise ValueError(
f"Population {population!r} has inconsistent prevalence values."
)
prevalence = next(iter(values))
thresholds = sorted(set(population_thresholds.get(population, [])))
references.append(
{
"type": "path",
"scope": "population",
"population": population,
"label": f"Treat All — {population}",
"benchmark": "treat_all",
"points": [
{
"x": threshold,
"y": prevalence
- (1.0 - prevalence) * threshold / (1.0 - threshold),
}
for threshold in thresholds
],
}
)

return {
"schemaVersion": "2.0",
"type": "decision_curve",
"evaluations": evaluations,
"series": series,
"data": data,
"x": "threshold",
"y": "netBenefit",
"xAxis": {
"label": "Probability threshold",
"domain": [min_p_threshold, max_p_threshold],
},
"yAxis": {"label": "Net benefit"},
"references": references,
}
1 change: 1 addition & 0 deletions src/rtichoke/_renderers.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ def write_html(self, path: str | Path) -> Path:
"precision_recall": "renderPrecisionRecallV2",
"gains": "renderGainsV2",
"lift": "renderLiftV2",
"decision_curve": "renderDecisionCurveV2",
}.get(str(self.spec.get("type")))
if render_export is None:
raise ValueError(
Expand Down
8 changes: 4 additions & 4 deletions src/rtichoke/_vendor/rtichoke_viz/VENDORED_FROM
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
repository=https://github.com/uriahf/rtichoke_viz
release=v0.5.0
source_commit=9c5a114ebe968e8cef4d2f14bf82ed552d2c8a17
archive=rtichoke-viz-0.5.0.tar.gz
sha256=ab36ae71f9090b4de62da8f552ebe84ac35885ab04958667faaf070db2c98f65
release=v0.6.0
source_commit=3abb3f07a598c3e22d5362a3f88e52bb6b52b083
archive=rtichoke-viz-0.6.0.tar.gz
sha256=625613c7f692ff50b7757a27bb6caf84e311971bde92593141393dbd897af3a2
Binary file not shown.
Binary file not shown.
Loading
Loading