Skip to content
Draft
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
14 changes: 4 additions & 10 deletions skrub/_data_ops/_data_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -747,21 +747,15 @@ def _repr_html_(self):
from ._inspection import node_report
from ._subsampling import uses_subsampling

graph_drawing = self.skb.draw_graph()
try:
graph = self.skb.draw_graph().svg.decode("utf-8")
graph = graph_drawing.svg.decode("utf-8")
graph = strip_xml_declaration(graph)
has_graph = True
except Exception:
graph = f"<p>{_utils.graphviz_error_message(html=True)}</p>"
has_graph = False
graph = graph_drawing.html_fragment
impl = self._skrub_impl
if impl.preview_if_available() is NULL:
if has_graph:
return f"<div>{graph}</div>"
return (
f"<div><div><strong><samp>{html.escape(short_repr(self))}</samp></strong>"
f"</div><div>{graph}</div></div>"
)
return f"<div>{graph}</div>"
if not isinstance(impl, Var) and impl.name is not None:
name_line = (
"<strong><samp>Name:"
Expand Down
66 changes: 53 additions & 13 deletions skrub/_data_ops/_inspection.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,19 @@
import base64
import copy
import datetime
import html
import io
import numbers
import re
import shutil
import sys
import uuid
import webbrowser
from pathlib import Path

import jinja2
import numpy as np
import pydot
from sklearn.base import BaseEstimator

from .. import _dataframe as sbd
Expand Down Expand Up @@ -38,6 +42,7 @@ def _get_jinja_env():
autoescape=True,
)
env.filters["format_duration"] = format_duration
env.globals["uuid"] = str(uuid.uuid4())
return env


Expand Down Expand Up @@ -145,7 +150,6 @@ def _make_full_report(
overwrite=False,
title=None,
):
_utils.check_graphviz()
output_dir = _get_output_dir(output_dir, overwrite)
try:
# TODO dump report in callback instead of evaluating full DataOps plan
Expand All @@ -166,7 +170,8 @@ def node_name_to_url(node_name):
def make_url(node):
return node_name_to_url(node_rindex[id(node)])

svg = draw_data_op_graph(data_op, url=make_url).svg.decode("utf-8")
graph_drawing = draw_data_op_graph(data_op, url=make_url)
svg = graph_drawing.html_fragment
jinja_env = _get_jinja_env()
index = jinja_env.get_template("index.html").render(
{"svg": svg, "node_status": node_status, "report_title": title}
Expand Down Expand Up @@ -220,6 +225,16 @@ def make_url(node):
estimator_html_repr = None
else:
estimator_html_repr = None
# TODO:
# - edit attributes to show node status (error, skipped)
# - edit instead of copy?
# - rename svg to graph_drawing_fragment
# - no need for highlightcurrentnode javascript function anymore
graph_drawing_copy = copy.deepcopy(graph_drawing)
dot_node = graph_drawing_copy.graph.get_node(_dot_id(i))[0]
dot_node.set_fillcolor("#15ed8f")
dot_node.set_style("filled")
svg = graph_drawing_copy.html_fragment
node_page = jinja_env.get_template("node.html").render(
dict(
report_title=title,
Expand Down Expand Up @@ -256,11 +271,20 @@ def make_url(node):


class GraphDrawing:
def __init__(self, graph):
def __init__(self, graph, force_js_rendering=False):
self.graph = graph
self.force_js_rendering = force_js_rendering

def _use_js(self):
return self.force_js_rendering or not _utils.has_graphviz()

def _base64(self):
dot = self.graph.to_string().encode("utf-8")
return base64.b64encode(dot).decode("ascii")

@property
def svg(self):
_utils.check_graphviz()
svg = self.graph.create_svg(encoding="utf-8")
svg = re.sub(b"<title>.*?</title>", b"", svg)
if "google.colab" in sys.modules:
Expand All @@ -271,15 +295,37 @@ def svg(self):

@property
def png(self):
_utils.check_graphviz()
return self.graph.create_png(encoding="utf-8")

def _repr_html_(self):
return self.svg.decode("utf-8")
if self._use_js():
return _get_template("render_dot_iframe.html").render(
{"dot_base64": self._base64()}
)
else:
return self.svg.decode("utf-8")

@property
def html_fragment(self):
if self._use_js():
return _get_template("render_dot_fragment.html").render(
{"dot_base64": self._base64()}
)
else:
return self.svg.decode("utf-8")

@property
def html(self):
if self._use_js():
return _get_template("render_dot.html").render(
{"dot_base64": self._base64()}
)
else:
return _get_template("graph.html").render({"svg": self.svg.decode("utf-8")})

def open(self):
open_in_browser(
_get_template("graph.html").render({"svg": self.svg.decode("utf-8")})
)
open_in_browser(self.html)

def _repr_png_(self):
return self.png
Expand Down Expand Up @@ -334,12 +380,6 @@ def _dot_id(n):


def draw_data_op_graph(data_op, *, url=None, direction="TB", show_ids=False):
# TODO if pydot or graphviz not available fallback on some other plotting
# solution eg a vendored copy of mermaid? outputting html instead of svg
_utils.check_graphviz()

import pydot

g = graph(data_op)
dot_graph = pydot.Dot(rankdir=direction, ranksep=0.4)
for node_id, e in g["nodes"].items():
Expand Down
6 changes: 2 additions & 4 deletions skrub/_data_ops/tests/test_inspection.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,15 +172,13 @@ def _import(name, *args, **kwargs):
return builtin_import(name, *args, **kwargs)

monkeypatch.setattr(builtins, "__import__", _import)
with pytest.raises(RuntimeError, match="please install Pydot and Graphviz"):
skrub.as_data_op(0).skb.draw_graph()
assert "Graphviz.load" in skrub.as_data_op(0).skb.draw_graph().html


def test_no_graphviz(monkeypatch):
pydot = pytest.importorskip("pydot")
monkeypatch.setattr(pydot.Dot, "create_svg", Mock(side_effect=Exception()))
with pytest.raises(RuntimeError, match="please install Pydot and Graphviz"):
skrub.as_data_op(0).skb.draw_graph()
assert "Graphviz.load" in skrub.as_data_op(0).skb.draw_graph().html


@pytest.mark.skipif(not _utils.has_graphviz(), reason="report requires graphviz")
Expand Down
21 changes: 15 additions & 6 deletions skrub/_data_ops/tests/test_interactive_features.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ def test_key_completions():
def test_repr_html():
a = skrub.var("thename", "thevalue")
r = a._repr_html_()
if "Please install" in r:
if "Graphviz.load" in r:
pytest.skip("graphviz not installed")
assert "thename" in r and "thevalue" in r
a = skrub.var("thename", skrub.datasets.toy_orders().orders)
Expand Down Expand Up @@ -354,29 +354,35 @@ def test_estimator_repr_html():
if old_sklearn:
assert data_op_repr in learner_repr
else:
assert svg in learner_repr
assert svg in learner_repr or data_op.skb.draw_graph()._base64() in learner_repr

X, y = make_classification()
learner.fit({"X": X, "y": y})
learner_repr = learner._repr_html_()
if old_sklearn:
assert data_op_repr in learner_repr
else:
assert svg in learner._repr_html_()
assert svg in learner_repr or data_op.skb.draw_graph()._base64() in learner_repr

rand_search_repr = data_op.skb.make_randomized_search()._repr_html_()
if old_sklearn:
assert data_op_repr in rand_search_repr
else:
assert svg in rand_search_repr
assert (
svg in rand_search_repr
or data_op.skb.draw_graph()._base64() in rand_search_repr
)
assert "n_iter" in rand_search_repr
assert "randomized search" in rand_search_repr

grid_search_repr = data_op.skb.make_grid_search()._repr_html_()
if old_sklearn:
assert data_op_repr in grid_search_repr
else:
assert svg in grid_search_repr
assert (
svg in rand_search_repr
or data_op.skb.draw_graph()._base64() in grid_search_repr
)
assert "n_jobs" in grid_search_repr
assert "grid search" in grid_search_repr

Expand All @@ -387,5 +393,8 @@ def test_estimator_repr_html():
if old_sklearn:
assert data_op_repr in optuna_search_repr
else:
assert svg in optuna_search_repr
assert (
svg in optuna_search_repr
or data_op.skb.draw_graph()._base64() in optuna_search_repr
)
assert "n_iter" in optuna_search_repr
10 changes: 10 additions & 0 deletions skrub/_reporting/_data/templates/data_ops/render_dot.html
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
<!DOCTYPE html>
<html>
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
</head>
<body>
{% include "render_dot_fragment.html" %}
</body>
</html>
12 changes: 12 additions & 0 deletions skrub/_reporting/_data/templates/data_ops/render_dot_fragment.html
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
<div id="{{ uuid }}">Loading graph …</div>
<script type="module">
import { Graphviz } from "https://cdn.jsdelivr.net/npm/@hpcc-js/wasm-graphviz@1/dist/index.js";
const dot = new TextDecoder().decode(Uint8Array.from(atob("{{ dot_base64 }}"), c => c.charCodeAt(0)));
try {
const graphviz = await Graphviz.load();
document.getElementById("{{ uuid }}").innerHTML = graphviz.dot(dot);
} catch (err) {
document.getElementById("{{ uuid }}").textContent = "Error: " + err.message;
}
parent.postMessage({ "{{ uuid }}_h": document.body.scrollHeight + 20, "{{ uuid }}_w": document.body.scrollWidth + 20}, "*");
</script>
13 changes: 13 additions & 0 deletions skrub/_reporting/_data/templates/data_ops/render_dot_iframe.html
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
<iframe srcdoc='{% include "render_dot.html" %}'
style="width:100%;border:none;"
id="{{ uuid }}">
</iframe>
<script>
window.addEventListener("message", e => {
if (e.data["{{ uuid }}_h"]) {
elem = document.getElementById("{{ uuid }}");
elem.style.height = e.data["{{ uuid }}_h"] + "px";
elem.style.width = e.data["{{ uuid }}_w"] + "px";
}
});
</script>
Loading