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
1 change: 1 addition & 0 deletions src/rtichoke/_report_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
"lift",
"decision_curve",
"interventions_avoided",
"prediction_distribution",
}

_ALL_SUPPORTED_TYPES = _V10_SCHEMA_TYPES | _V20_SCHEMA_TYPES
Expand Down
48 changes: 46 additions & 2 deletions src/rtichoke/summary_report/summary_report.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
_lift_v2_spec_from_performance_data,
_precision_recall_times_v2_spec_from_performance_data,
_precision_recall_v2_spec_from_performance_data,
_prediction_distribution_v2_spec,
_roc_times_v2_spec_from_performance_data,
_roc_v2_spec_from_performance_data,
)
Expand All @@ -43,6 +44,9 @@
_create_calibration_curve_list_times,
)
from rtichoke.performance_data.performance_data import prepare_performance_data
from rtichoke.performance_data.probs_distribution import (
_prepare_probs_distribution_data,
)
from rtichoke.performance_data.performance_data_times import (
prepare_performance_data_times,
)
Expand Down Expand Up @@ -328,6 +332,9 @@ def create_summary_report(
is an explicit opt-in path that uses Python's existing production
calculations, canonical standalone component builders, canonical ReportSpec
assembly, and the vendored ``rtichoke_viz`` ``renderReport()`` composer.
For ``renderer="browser"``, the static Discrimination section includes
Prediction Distribution by Probability Threshold and Prediction Distribution
by PPCR / Risk Percentile.

Parameters
----------
Expand Down Expand Up @@ -371,15 +378,42 @@ def _create_browser_summary_report(
output_file: str | Path,
) -> Path:
"""Build the canonical static ReportSpec v1.1 browser summary report."""
by = 0.01
metadata = _build_evaluation_metadata(probs, reals, np.array([]))

# Stratified by probability threshold
perf_data_thresh = prepare_performance_data(
probs, reals, stratified_by=("probability_threshold",), by=0.01
probs, reals, stratified_by=("probability_threshold",), by=by
)
# Stratified by PPCR
perf_data_ppcr = prepare_performance_data(
probs, reals, stratified_by=("ppcr",), by=0.01
probs, reals, stratified_by=("ppcr",), by=by
)

threshold_distribution_data = _prepare_probs_distribution_data(
probs=probs,
reals=reals,
by=by,
stratified_by=("probability_threshold",),
)
ppcr_distribution_data = _prepare_probs_distribution_data(
probs=probs,
reals=reals,
by=by,
stratified_by=("ppcr",),
)

threshold_prediction_distribution = _prediction_distribution_v2_spec(
distribution_data=threshold_distribution_data,
performance_data=perf_data_thresh,
evaluation_metadata=metadata,
stratified_by=("probability_threshold",),
)
ppcr_prediction_distribution = _prediction_distribution_v2_spec(
distribution_data=ppcr_distribution_data,
performance_data=perf_data_ppcr,
evaluation_metadata=metadata,
stratified_by=("ppcr",),
)

calibration_curve_list = _create_calibration_curve_list(probs, reals)
Expand Down Expand Up @@ -478,6 +512,11 @@ def _create_browser_summary_report(
"id": "discrimination-probability-threshold",
"title": "By Probability Threshold",
"components": [
{
"id": "prediction-distribution",
"title": "Prediction Distribution",
"spec": threshold_prediction_distribution,
},
{"id": "roc", "title": "ROC", "spec": roc_thresh_spec},
{
"id": "precision-recall",
Expand All @@ -496,6 +535,11 @@ def _create_browser_summary_report(
"id": "discrimination-ppcr",
"title": "By PPCR",
"components": [
{
"id": "prediction-distribution-2",
"title": "Prediction Distribution",
"spec": ppcr_prediction_distribution,
},
{"id": "roc-2", "title": "ROC", "spec": roc_ppcr_spec},
{
"id": "precision-recall-2",
Expand Down
2 changes: 2 additions & 0 deletions tests/test_quarto_summary_report_browser.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,7 @@ def test_quarto_single_browser_summary_report(tmp_path):
"Model" in tbl_text
or "Probability Threshold" in tbl_text
or "True Positives" in tbl_text
or "Sensitivity" in tbl_text
)
assert frame.locator("svg").count() >= 2
assert len(errors) == 0, f"Console errors found: {errors}"
Expand Down Expand Up @@ -243,6 +244,7 @@ def test_quarto_two_browser_summary_reports(tmp_path):
"Model" in tbl_text
or "Probability Threshold" in tbl_text
or "True Positives" in tbl_text
or "Sensitivity" in tbl_text
)
assert frame.locator("svg").count() >= 2

Expand Down
23 changes: 23 additions & 0 deletions tests/test_report_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,29 @@ def test_type_aware_schema_version_validation() -> None:
_build_report_spec_v11(sections_v2)


def test_prediction_distribution_schema_version_validation() -> None:
valid_pred_dist = _curve_spec("prediction_distribution")
sections_valid = [
{
"id": "discrimination",
"components": [{"id": "prediction-distribution", "spec": valid_pred_dist}],
}
]
report = _build_report_spec_v11(sections_valid)
assert report["schemaVersion"] == "1.1"

bad_pred_dist = _curve_spec("prediction_distribution")
bad_pred_dist["schemaVersion"] = "1.0"
sections_invalid = [
{
"id": "discrimination",
"components": [{"id": "prediction-distribution", "spec": bad_pred_dist}],
}
]
with pytest.raises(ValueError, match="requires schemaVersion '2.0'"):
_build_report_spec_v11(sections_invalid)


def test_existing_summary_report_still_uses_r_backend(
monkeypatch: pytest.MonkeyPatch,
) -> None:
Expand Down
Loading
Loading