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
9 changes: 8 additions & 1 deletion .github/workflows/docs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,7 @@ jobs:
grep -q '.Reactable {' performance-table-reactable.html
grep -q 'Real Positive' performance-table-reactable.html
mv performance-table-reactable.html great-docs/_site/performance-table-reactable.html
cp -R examples/performance_table_reactable_files great-docs/_site/performance-table-reactable_files
cp -R examples/performance_table_reactable_files great-docs/_site/performance_table_reactable_files

- name: Generate rtichoke_viz ROC proof
if: github.event.action != 'closed'
Expand All @@ -102,6 +102,13 @@ jobs:
mkdir -p great-docs/_site/rtichoke-viz-roc
cp -R rtichoke_viz_roc_proof/. great-docs/_site/rtichoke-viz-roc/

- name: Generate rtichoke_viz calibration proof
if: github.event.action != 'closed'
run: |
uv run python examples/rtichoke_viz_calibration.py
mkdir -p great-docs/_site/rtichoke-viz-calibration
cp -R rtichoke_viz_calibration_proof/. great-docs/_site/rtichoke-viz-calibration/

- name: Deploy PR preview
uses: rossjrw/pr-preview-action@v1
with:
Expand Down
22 changes: 22 additions & 0 deletions examples/rtichoke_viz_calibration.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
"""Generate a browser-rendered calibration proof from real rtichoke output."""

from pathlib import Path

import numpy as np

from rtichoke._viz_browser import _write_calibration_browser_html
from rtichoke.calibration.calibration import _create_calibration_curve_list

probs = {
"Model A": np.array(
[0.03, 0.08, 0.12, 0.18, 0.25, 0.32, 0.40, 0.50, 0.62, 0.75, 0.88, 0.96]
)
}
reals = np.array([0, 0, 0, 0, 0, 1, 0, 1, 1, 1, 1, 1])

calibration_curve_list = _create_calibration_curve_list(probs, reals)
output = _write_calibration_browser_html(
calibration_curve_list,
Path("rtichoke_viz_calibration_proof") / "index.html",
)
print(output)
95 changes: 89 additions & 6 deletions src/rtichoke/_viz_browser.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import json
from importlib.resources import files
from pathlib import Path
from typing import Any

import polars as pl

Expand All @@ -16,6 +17,16 @@
}


def _copy_viz_assets(output_dir: Path) -> None:
vendor = files("rtichoke").joinpath("_vendor", "rtichoke_viz")
(output_dir / "rtichoke-viz.js").write_bytes(
vendor.joinpath("rtichoke-viz.js").read_bytes()
)
(output_dir / "rtichoke-viz.css").write_bytes(
vendor.joinpath("rtichoke-viz.css").read_bytes()
)


def _roc_spec_from_performance_data(
performance_data: pl.DataFrame,
) -> dict[str, object]:
Expand Down Expand Up @@ -52,19 +63,56 @@ def _roc_spec_from_performance_data(
}


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(
"reference_group", "x", "y", "n_reals", "n"
)
distribution_rows = calibration_curve_list["histogram_for_calibration"].select(
"reference_group", "mids", "counts"
)

return {
"schemaVersion": "1.0",
"type": "calibration",
"data": [
{
"model": row["reference_group"],
"predicted": row["x"],
"observed": row["y"],
"method": "discrete",
"events": row["n_reals"],
"total": row["n"],
}
for row in rows.to_dicts()
],
"distribution": [
{
"model": row["reference_group"],
"midpoint": row["mids"],
"count": row["counts"],
"binWidth": 0.01,
}
for row in distribution_rows.to_dicts()
],
"x": "predicted",
"y": "observed",
"xAxis": {"label": "Predicted probability", "domain": [0, 1]},
"yAxis": {"label": "Observed probability", "domain": [0, 1]},
"references": [{"type": "identity"}],
}


def _write_roc_browser_html(
performance_data: pl.DataFrame,
output_path: str | Path,
) -> Path:
"""Write a standalone proof page that uses the vendored browser renderer."""
output = Path(output_path)
output.parent.mkdir(parents=True, exist_ok=True)

vendor = files("rtichoke").joinpath("_vendor", "rtichoke_viz")
js_source = vendor.joinpath("rtichoke-viz.js")
css_source = vendor.joinpath("rtichoke-viz.css")
(output.parent / "rtichoke-viz.js").write_bytes(js_source.read_bytes())
(output.parent / "rtichoke-viz.css").write_bytes(css_source.read_bytes())
_copy_viz_assets(output.parent)

spec_json = json.dumps(_roc_spec_from_performance_data(performance_data)).replace(
"</", "<\\/"
Expand All @@ -90,3 +138,38 @@ def _write_roc_browser_html(
"""
output.write_text(html, encoding="utf-8")
return output


def _write_calibration_browser_html(
calibration_curve_list: dict[str, Any],
output_path: str | Path,
) -> Path:
"""Write a standalone discrete-calibration proof using vendored assets."""
output = Path(output_path)
output.parent.mkdir(parents=True, exist_ok=True)
_copy_viz_assets(output.parent)

spec_json = json.dumps(
_calibration_spec_from_curve_list(calibration_curve_list)
).replace("</", "<\\/")
html = f"""<!doctype html>
<html lang=\"en\">
<head>
<meta charset=\"utf-8\">
<meta name=\"viewport\" content=\"width=device-width, initial-scale=1\">
<title>rtichoke_viz calibration vendoring proof</title>
<link rel=\"stylesheet\" href=\"./rtichoke-viz.css\">
</head>
<body>
<div id=\"calibration-chart\" class=\"rtichoke-viz-chart\"></div>
<script id=\"calibration-spec\" type=\"application/json\">{spec_json}</script>
<script type=\"module\">
import {{ renderCalibration }} from \"./rtichoke-viz.js\";
const spec = JSON.parse(document.querySelector(\"#calibration-spec\").textContent);
document.querySelector(\"#calibration-chart\").append(renderCalibration(spec));
</script>
</body>
</html>
"""
output.write_text(html, encoding="utf-8")
return output
43 changes: 42 additions & 1 deletion tests/test_viz_browser.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,12 @@
import numpy as np

from rtichoke._viz_browser import (
_calibration_spec_from_curve_list,
_roc_spec_from_performance_data,
_write_calibration_browser_html,
_write_roc_browser_html,
)
from rtichoke.calibration.calibration import _create_calibration_curve_list
from rtichoke.performance_data.performance_data import prepare_performance_data


Expand All @@ -17,6 +20,17 @@ def _real_roc_performance_data():
)


def _real_calibration_curve_list():
return _create_calibration_curve_list(
probs={
"Model A": np.array(
[0.03, 0.08, 0.12, 0.18, 0.25, 0.32, 0.40, 0.50, 0.62, 0.75, 0.88, 0.96]
)
},
reals=np.array([0, 0, 0, 0, 0, 1, 0, 1, 1, 1, 1, 1]),
)


def test_real_roc_output_maps_to_canonical_spec():
performance_data = _real_roc_performance_data()

Expand All @@ -31,7 +45,21 @@ def test_real_roc_output_maps_to_canonical_spec():
assert all(0 <= row["specificity"] <= 1 for row in spec["data"])


def test_browser_proof_uses_vendored_assets(tmp_path: Path):
def test_real_calibration_output_maps_to_canonical_spec():
spec = _calibration_spec_from_curve_list(_real_calibration_curve_list())

assert spec["schemaVersion"] == "1.0"
assert spec["type"] == "calibration"
assert spec["references"] == [{"type": "identity"}]
assert spec["data"]
assert {row["model"] for row in spec["data"]} == {"Model A"}
assert all(row["method"] == "discrete" for row in spec["data"])
assert all(0 <= row["predicted"] <= 1 for row in spec["data"])
assert all(0 <= row["observed"] <= 1 for row in spec["data"])
assert spec["distribution"]


def test_roc_browser_proof_uses_vendored_assets(tmp_path: Path):
output = _write_roc_browser_html(
_real_roc_performance_data(),
tmp_path / "index.html",
Expand All @@ -42,3 +70,16 @@ def test_browser_proof_uses_vendored_assets(tmp_path: Path):
assert '"type": "roc"' in html
assert (tmp_path / "rtichoke-viz.js").stat().st_size > 0
assert (tmp_path / "rtichoke-viz.css").stat().st_size > 0


def test_calibration_browser_proof_uses_vendored_assets(tmp_path: Path):
output = _write_calibration_browser_html(
_real_calibration_curve_list(),
tmp_path / "index.html",
)

html = output.read_text(encoding="utf-8")
assert 'import { renderCalibration } from "./rtichoke-viz.js"' in html
assert '"type": "calibration"' in html
assert (tmp_path / "rtichoke-viz.js").stat().st_size > 0
assert (tmp_path / "rtichoke-viz.css").stat().st_size > 0
Loading