diff --git a/.github/workflows/python-package.yml b/.github/workflows/python-package.yml index eea114ee..7aa6989f 100644 --- a/.github/workflows/python-package.yml +++ b/.github/workflows/python-package.yml @@ -52,7 +52,7 @@ jobs: prefix = "rtichoke/_vendor/rtichoke_viz/" required = { f"{prefix}VENDORED_FROM", - f"{prefix}rtichoke-viz-0.20.2.tar.gz", + f"{prefix}rtichoke-viz-0.22.1.tar.gz", f"{prefix}rtichoke-viz.js", f"{prefix}rtichoke-viz.css", f"{prefix}rtichoke-viz.schema.json", @@ -60,6 +60,8 @@ jobs: f"{prefix}rtichoke-viz-report.schema.json", } assert required <= names + assert f"{prefix}rtichoke-viz-0.22.0.tar.gz" not in names + assert f"{prefix}rtichoke-viz-0.20.2.tar.gz" not in names assert f"{prefix}rtichoke-viz-0.20.1.tar.gz" not in names assert f"{prefix}rtichoke-viz-0.20.0.tar.gz" not in names assert f"{prefix}rtichoke-viz-0.19.0.tar.gz" not in names diff --git a/.github/workflows/quarto-acceptance.yml b/.github/workflows/quarto-acceptance.yml index 2ef7904a..33344b05 100644 --- a/.github/workflows/quarto-acceptance.yml +++ b/.github/workflows/quarto-acceptance.yml @@ -15,6 +15,8 @@ on: - "tests/test_summary_report_browser.py" - "tests/test_summary_report_times_browser.py" - "tests/test_decision_curve_browser_acceptance.py" + - "tests/test_probs_histogram_browser_acceptance.py" + - "src/rtichoke/probs_distribution.py" - ".github/workflows/quarto-acceptance.yml" pull_request: branches: ["main"] @@ -30,6 +32,8 @@ on: - "tests/test_summary_report_browser.py" - "tests/test_summary_report_times_browser.py" - "tests/test_decision_curve_browser_acceptance.py" + - "tests/test_probs_histogram_browser_acceptance.py" + - "src/rtichoke/probs_distribution.py" - ".github/workflows/quarto-acceptance.yml" workflow_dispatch: @@ -65,3 +69,4 @@ jobs: tests/test_decision_curve_browser_acceptance.py tests/test_summary_report_browser.py tests/test_summary_report_times_browser.py + tests/test_probs_histogram_browser_acceptance.py diff --git a/src/rtichoke/__init__.py b/src/rtichoke/__init__.py index 896d7eb1..ff79f37f 100644 --- a/src/rtichoke/__init__.py +++ b/src/rtichoke/__init__.py @@ -57,12 +57,17 @@ render_performance_table as render_performance_table, ) +from rtichoke.probs_distribution import ( + create_probs_histogram as create_probs_histogram, +) + from rtichoke.summary_report.summary_report import ( create_summary_report as create_summary_report, create_summary_report_times as create_summary_report_times, ) __all__ = [ + "create_probs_histogram", "create_roc_curve", "create_roc_curve_times", "plot_roc_curve", diff --git a/src/rtichoke/_renderers.py b/src/rtichoke/_renderers.py index fae719e0..2d10c95f 100644 --- a/src/rtichoke/_renderers.py +++ b/src/rtichoke/_renderers.py @@ -8,6 +8,8 @@ from pathlib import Path from typing import Any, Literal +from rtichoke._report_browser import _resolve_render_report_symbol, _sanitize_nan_values + Renderer = Literal["plotly", "matplotlib", "browser", "rtichoke_viz"] _SUPPORTED_RENDERERS = ("plotly", "matplotlib", "browser", "rtichoke_viz") @@ -46,31 +48,69 @@ def write_html(self, path: str | Path) -> Path: output = Path(path) output.parent.mkdir(parents=True, exist_ok=True) vendor = files("rtichoke").joinpath("_vendor", "rtichoke_viz") - for asset in ("rtichoke-viz.js", "rtichoke-viz.css"): - (output.parent / asset).write_bytes(vendor.joinpath(asset).read_bytes()) - render_export = { - "roc": "renderRocV2", - "calibration": "renderCalibrationV2", - "precision_recall": "renderPrecisionRecallV2", - "gains": "renderGainsV2", - "lift": "renderLiftV2", - "decision_curve": "renderDecisionCurveV2", - "interventions_avoided": "renderInterventionsAvoidedV2", - }.get(str(self.spec.get("type"))) - if render_export is None: - raise ValueError( - f"rtichoke_viz does not support chart type {self.spec.get('type')!r}." + chart_type = str(self.spec.get("type")) + if chart_type == "prediction_distribution": + viz_js = vendor.joinpath("rtichoke-viz.js").read_text(encoding="utf-8") + viz_css = vendor.joinpath("rtichoke-viz.css").read_text(encoding="utf-8") + render_fn = _resolve_render_report_symbol( + viz_js, "renderPredictionDistribution" + ) + sanitized_spec = _sanitize_nan_values(self.spec) + spec_json = json.dumps(sanitized_spec, separators=(",", ":")).replace( + "", "<\\/" ) - spec_json = json.dumps(self.spec, separators=(",", ":")).replace("", "<\\/") - html = f""" + html = f""" + +
+ + + +