Static and interactive visualizations of fitted models.
draw() renders a model as a self-contained SVG figure (save it as .svg, .png, .pdf or .html), and
interactive() builds a single offline HTML page to explore it: fold and unfold trees, highlight a
feature, route a sample through the model, and see the smallest change that would flip its prediction.
Supported models, each drawn so that it reproduces the model's own predict / predict_proba:
- imodels: trees (GreedyTree, DecisionTreeCCP, HSTree, TaoTree, C4.5, FastSmallTree), sums of trees (FIGS, IRF), rule lists (GreedyRuleList, OneR, FastFrugalTree, BayesianRuleList), rule sets (RuleFit, FPLasso, SkopeRules, FPSkope, BoostedRules, Slipper, BayesianRuleSet), scoring systems (FastRiskScore, SLIM) and additive models (TreeGAM, GPGam, MarginalShrinkageLinearRegressor).
- scikit-learn: decision trees, random forests, extra trees, gradient boosting, histogram gradient boosting, linear and logistic models, GLMs, linear SVMs and isotonic regression. Pipelines whose earlier steps only scale features are drawn in raw units.
from imodels import FIGSClassifier
from imodels.viz import draw, interactive
model = FIGSClassifier(max_rules=8).fit(X, y)
draw(model, X, y).save("figs.svg") # static figure
interactive(model, X, y).save("figs.html") # interactive page, one offline file
Both functions show inline in Jupyter. See the gallery post for examples of every kind of model.
Tree-based imodels models (the trees above, FIGS and IRF) are drawn through their exact scikit-learn
export, to_sklearn(), so this package only reads scikit-learn trees. The same export makes
every tree-based imodels model drawable by any tool that reads scikit-learn trees, such as
dtreeviz (dtreeviz.model(to_sklearn()(model), X, y)) or
sklearn.tree.plot_tree.
Expand source code
"""Static and interactive visualizations of fitted models.
`draw` renders a model as a self-contained SVG figure (save it as .svg, .png, .pdf or .html), and
`interactive` builds a single offline HTML page to explore it: fold and unfold trees, highlight a
feature, route a sample through the model, and see the smallest change that would flip its prediction.
Supported models, each drawn so that it reproduces the model's own `predict` / `predict_proba`:
- **imodels**: trees (GreedyTree, DecisionTreeCCP, HSTree, TaoTree, C4.5, FastSmallTree), sums of trees
(FIGS, IRF), rule lists (GreedyRuleList, OneR, FastFrugalTree, BayesianRuleList), rule sets (RuleFit,
FPLasso, SkopeRules, FPSkope, BoostedRules, Slipper, BayesianRuleSet), scoring systems
(FastRiskScore, SLIM) and additive models (TreeGAM, GPGam, MarginalShrinkageLinearRegressor).
- **scikit-learn**: decision trees, random forests, extra trees, gradient boosting, histogram gradient
boosting, linear and logistic models, GLMs, linear SVMs and isotonic regression. Pipelines whose
earlier steps only scale features are drawn in raw units.
```python
from imodels import FIGSClassifier
from imodels.viz import draw, interactive
model = FIGSClassifier(max_rules=8).fit(X, y)
draw(model, X, y).save("figs.svg") # static figure
interactive(model, X, y).save("figs.html") # interactive page, one offline file
```
Both functions show inline in Jupyter. See the [gallery post](https://csinva.io/imodels/viz.html)
for examples of every kind of model.
Tree-based imodels models (the trees above, FIGS and IRF) are drawn through their exact scikit-learn
export, `imodels.to_sklearn`, so this package only reads scikit-learn trees. The same export makes
every tree-based imodels model drawable by any tool that reads scikit-learn trees, such as
[dtreeviz](https://github.com/parrt/dtreeviz) (`dtreeviz.model(imodels.to_sklearn(model), X, y)`) or
`sklearn.tree.plot_tree`.
"""
from ._interactive import InteractiveTree, interactive
from ._static import TreeFigure, draw
from ._textview import text
__all__ = ["draw", "interactive", "text", "TreeFigure", "InteractiveTree"]
Functions
def draw(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, x=None, max_depth=None, orientation='auto', theme='light', style='auto', title=None, subtitle=None, legend=True, background=True, precision=3, simple=False, output=0, max_trees=None)-
Draw a fitted model as a static figure.
Returns a
TreeFigure, which shows inline in Jupyter and saves to .svg, .png, .pdf or .html (TreeFigure.save()). PNG and PDF need cairosvg.Parameters
model:estimator- A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features.
X:array-likeofshape (n_samples, n_features), optional- Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally.
y:array-likeofshape (n_samples,), optional- Targets for
X, used for class mixes and target ranges. feature_names:listofstr, optional- Feature names. Default: the columns of
X, else the names the model was fitted with. class_names:listordict, optional- Class names, as a list or as a dict from class label to name. Default: the model's classes.
target_name:str, optional- Name of the target (regression). Default: the name of
y, else "target". x:array-likeofshape (n_features,), optional- A single sample whose path through the model is highlighted.
max_depth:int, optional- Draw only the top
max_depthlevels; deeper subtrees are drawn as stacked cards. orientation:{"auto", "LR", "TB"}, default="auto"- Left to right, or top to bottom. "auto" lays a single tree out left to right and several trees (forests, boosting, FIGS) top to bottom.
theme:{"light", "dark"}, default="light"- Color theme.
style:{"auto", "detailed", "compact"}, default="auto"- Card style; "auto" is compact for trees with more than 24 leaves.
title,subtitle:str, optional- Figure title, and a subtitle (default: a one-line summary of the model).
legend:bool, default=True- Whether to draw the legend.
background:bool, default=True- Whether to fill the background (False gives a transparent figure).
precision:int, default=3- Significant digits for thresholds and values.
simple:bool, default=False- Draw each node as a plain box filled by its class mix (or predicted value) instead of the detailed card with a chart.
output:int, default=0- For multi-output trees, which output to draw.
max_trees:int, optional- For ensembles (forests, boosting, FIGS), how many trees to draw (default 6). Predictions always use every tree.
Returns
TreeFigure- The figure, with
.svg(the SVG text) and.save(path).
Expand source code
def draw(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, x=None, max_depth=None, orientation="auto", theme="light", style="auto", title=None, subtitle=None, legend=True, background=True, precision=3, simple=False, output=0, max_trees=None): """Draw a fitted model as a static figure. Returns a `TreeFigure`, which shows inline in Jupyter and saves to .svg, .png, .pdf or .html (`TreeFigure.save`). PNG and PDF need cairosvg. Parameters ---------- model : estimator A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features. X : array-like of shape (n_samples, n_features), optional Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally. y : array-like of shape (n_samples,), optional Targets for ``X``, used for class mixes and target ranges. feature_names : list of str, optional Feature names. Default: the columns of ``X``, else the names the model was fitted with. class_names : list or dict, optional Class names, as a list or as a dict from class label to name. Default: the model's classes. target_name : str, optional Name of the target (regression). Default: the name of ``y``, else "target". x : array-like of shape (n_features,), optional A single sample whose path through the model is highlighted. max_depth : int, optional Draw only the top ``max_depth`` levels; deeper subtrees are drawn as stacked cards. orientation : {"auto", "LR", "TB"}, default="auto" Left to right, or top to bottom. "auto" lays a single tree out left to right and several trees (forests, boosting, FIGS) top to bottom. theme : {"light", "dark"}, default="light" Color theme. style : {"auto", "detailed", "compact"}, default="auto" Card style; "auto" is compact for trees with more than 24 leaves. title, subtitle : str, optional Figure title, and a subtitle (default: a one-line summary of the model). legend : bool, default=True Whether to draw the legend. background : bool, default=True Whether to fill the background (False gives a transparent figure). precision : int, default=3 Significant digits for thresholds and values. simple : bool, default=False Draw each node as a plain box filled by its class mix (or predicted value) instead of the detailed card with a chart. output : int, default=0 For multi-output trees, which output to draw. max_trees : int, optional For ensembles (forests, boosting, FIGS), how many trees to draw (default 6). Predictions always use every tree. Returns ------- TreeFigure The figure, with ``.svg`` (the SVG text) and ``.save(path)``. """ info = to_view(model, X, y, feature_names, class_names, target_name, output) if isinstance(info, AdditiveView): from ._render_views import render_view return TreeFigure(render_view(info, Paint(theme), precision, new_uid(), x=x, title=title, subtitle=subtitle)) if max_trees is not None: info.shown_roots = info.roots[:max(1, int(max_trees))] style = _resolve_style(style, info, max_depth) P = Paint(theme) artist = Artist(info, P, style=style, sig=precision, uid=new_uid(), x=x, simple=simple) pred = None path = [] if x is not None: path = instance_path(info, x) pred = prediction_text(info, path[-1], precision) artist.set_path(path) fig = Figure(info, artist, P, orientation=resolve_orientation(orientation, info), max_depth=max_depth, title=title, subtitle=subtitle, legend=legend, background=background, prediction=pred) return TreeFigure(fig.svg()) def interactive(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, title=None, subtitle=None, theme='light', orientation='auto', style='auto', initial_depth=None, precision=3, max_samples=400, simple='auto', output=0, height=720, max_trees=None)-
Build an interactive page for a fitted model: one self-contained HTML file.
Click a split to fold it, drag to pan, scroll to zoom, and hover for the full rule. The Predict panel routes a typed or sampled input through the model, shows how its prediction is built, and lists the smallest changes that would flip it. The Features panel shows each feature's importance and highlights where it is used. Shows inline in Jupyter; save with
.save(path).Parameters
model:estimator- A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features.
X:array-likeofshape (n_samples, n_features), optional- Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally.
y:array-likeofshape (n_samples,), optional- Targets for
X, used for class mixes and target ranges. feature_names:listofstr, optional- Feature names. Default: the columns of
X, else the names the model was fitted with. class_names:listordict, optional- Class names, as a list or as a dict from class label to name. Default: the model's classes.
target_name:str, optional- Name of the target (regression). Default: the name of
y, else "target". title,subtitle:str, optional- Page title, and a subtitle (default: a one-line summary of the model).
theme:{"light", "dark", "auto"}, default="light"- Color theme; "auto" follows the viewer's OS. The page has a toggle either way.
orientation:{"auto", "LR", "TB"}, default="auto"- Left to right, or top to bottom. "auto" lays a single tree out left to right and several trees (forests, boosting, FIGS) top to bottom.
style:{"auto", "detailed", "compact"}, default="auto"- Card style for trees.
initial_depth:int, optional- Levels expanded on load (default: all if the tree has at most 63 nodes, else 3).
precision:int, default=3- Significant digits for thresholds and values.
max_samples:int, default=400- Rows of
Xembedded for the "random sample" button (0 embeds none). simple:boolor"auto", default="auto"- Start in simple mode (plain boxes filled by class mix); "auto" does so for trees with more than 32 leaves. The page has a toggle either way.
output:int, default=0- For multi-output trees, which output to show.
height:int, default=720- Height in pixels when shown inline in Jupyter.
max_trees:int, optional- For ensembles (forests, boosting, FIGS), how many trees to draw (default 6). Predictions always use every tree.
Returns
InteractiveTree- The page, with
.html(the HTML text) and.save(path).
Expand source code
def interactive(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, title=None, subtitle=None, theme="light", orientation="auto", style="auto", initial_depth=None, precision=3, max_samples=400, simple="auto", output=0, height=720, max_trees=None): """Build an interactive page for a fitted model: one self-contained HTML file. Click a split to fold it, drag to pan, scroll to zoom, and hover for the full rule. The Predict panel routes a typed or sampled input through the model, shows how its prediction is built, and lists the smallest changes that would flip it. The Features panel shows each feature's importance and highlights where it is used. Shows inline in Jupyter; save with ``.save(path)``. Parameters ---------- model : estimator A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features. X : array-like of shape (n_samples, n_features), optional Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally. y : array-like of shape (n_samples,), optional Targets for ``X``, used for class mixes and target ranges. feature_names : list of str, optional Feature names. Default: the columns of ``X``, else the names the model was fitted with. class_names : list or dict, optional Class names, as a list or as a dict from class label to name. Default: the model's classes. target_name : str, optional Name of the target (regression). Default: the name of ``y``, else "target". title, subtitle : str, optional Page title, and a subtitle (default: a one-line summary of the model). theme : {"light", "dark", "auto"}, default="light" Color theme; "auto" follows the viewer's OS. The page has a toggle either way. orientation : {"auto", "LR", "TB"}, default="auto" Left to right, or top to bottom. "auto" lays a single tree out left to right and several trees (forests, boosting, FIGS) top to bottom. style : {"auto", "detailed", "compact"}, default="auto" Card style for trees. initial_depth : int, optional Levels expanded on load (default: all if the tree has at most 63 nodes, else 3). precision : int, default=3 Significant digits for thresholds and values. max_samples : int, default=400 Rows of ``X`` embedded for the "random sample" button (0 embeds none). simple : bool or "auto", default="auto" Start in simple mode (plain boxes filled by class mix); "auto" does so for trees with more than 32 leaves. The page has a toggle either way. output : int, default=0 For multi-output trees, which output to show. height : int, default=720 Height in pixels when shown inline in Jupyter. max_trees : int, optional For ensembles (forests, boosting, FIGS), how many trees to draw (default 6). Predictions always use every tree. Returns ------- InteractiveTree The page, with ``.html`` (the HTML text) and ``.save(path)``. """ info = to_view(model, X, y, feature_names, class_names, target_name, output) from ._views import AdditiveView if isinstance(info, AdditiveView): from ._interactive_view import interactive_view return interactive_view(info, title=title, subtitle=subtitle, theme=theme, precision=precision, max_samples=max_samples, height=height) style = _resolve_style(style, info, None) if style != "auto" or info.X is None else ( "detailed" if info.n_leaves <= 300 else "compact") P = Paint("light", use_vars=True) uid = new_uid() # big trees embed every card, so keep the per-node scatter lighter A = Artist(info, P, style=style, sig=precision, uid=uid, max_points=260 if len(info.nodes) <= 63 else 110) A.set_path([]) if max_trees is not None: info.shown_roots = info.roots[:max(1, int(max_trees))] shown = set(info.shown_roots or info.roots) root_of = {} for r in info.roots: stack = [r] while stack: i = stack.pop() root_of[i] = r if not info.nodes[i].is_leaf: stack += [info.nodes[i].left, info.nodes[i].right] nodes = [] for nd in info.nodes: lab = info.edge_label(nd.id, precision) if root_of.get(nd.id) not in shown: # a tree that is not drawn: only what prediction needs nodes.append(dict( id=nd.id, p=nd.parent, l=nd.left, r=nd.right, d=nd.depth, hidden=True, leaf=nd.is_leaf, c=[[int(f), op, float(v)] for f, op, v in nd.split], fs=nd.features, f=nd.feature if nd.simple else -1, nl=nd.nan_left, n=nd.n, wt=nd.weight, lab=lab, v=getattr(nd, "value", None) if getattr(nd, "value", None) is not None else (float(nd.counts[0]) if not info.is_clf else None), pr=[round(float(q), 6) for q in nd.counts / (nd.counts.sum() or 1)] if info.is_clf else None, pc=info.prediction(nd) if info.is_clf else None)) continue c = A.card(nd.id) sc = A.card(nd.id, simple=True) nodes.append(dict(nl=nd.nan_left, id=nd.id, p=nd.parent, l=nd.left, r=nd.right, d=nd.depth, w=round(c["w"], 1), h=round(c["h"], 1), svg=c["svg"], sw=round(sc["w"], 1), sh=round(sc["h"], 1), ssvg=sc["svg"], mix=[[col, round(f, 5)] for col, f in A.mix(nd)], n=nd.n, wt=nd.weight, col=A.node_color(nd), lab=lab, labw=round(text_width(lab, 10.5, True) + 14, 1), f=nd.feature if nd.simple else -1, t=nd.threshold, c=[[int(f), op, float(v)] for f, op, v in nd.split], fs=nd.features, st=info.split_text(nd.id, precision) if not nd.is_leaf else "", desc=info.n_descendants(nd.id), leaf=nd.is_leaf, chart=c["chart"], v=getattr(nd, "value", None) if getattr(nd, "value", None) is not None else (float(nd.counts[0]) if not info.is_clf else None), tip=_tooltip(info, A, nd.id, precision), pr=[round(float(q), 6) for q in nd.counts / (nd.counts.sum() or 1)] if info.is_clf else None, pc=info.prediction(nd) if info.is_clf else None, )) if info.layout == "cascade": # name each rule's outcome after its rule for nd in info.nodes: if not nd.is_leaf and nd.label: side = nd.left if info.nodes[nd.left].is_leaf else nd.right nodes[side]["ruleLabel"] = f"{nd.label} outcome" imp, imp_kind = info.importances() if not isinstance(model, TreeInfo) and info.combine == "single" and info.layout == "tree": try: # sklearn trees: use the estimator's own numbers imp, imp_kind = _unwrap(model).feature_importances_, "impurity" except (TypeError, AttributeError): pass used = sorted({f for nd in info.nodes for f in nd.features}, key=lambda f: -imp[f]) feats = [] for f in used: thr = [v for nd in info.nodes for g, _, v in nd.split if g == f] if info.X is not None: col = info.X[:, f] col = col[np.isfinite(col)] lo, hi, med = float(col.min()), float(col.max()), float(np.median(col)) else: lo, hi = min(thr), max(thr) span = (hi - lo) or abs(hi) or 1.0 lo, hi, med = lo - 0.25 * span, hi + 0.25 * span, float(np.median(thr)) kind = info.feature_kind.get(f, "continuous") step = 1 if kind in ("binary", "integer") else float(f"{(hi - lo) / 200:.2g}") or 0.01 feats.append(dict(i=f, name=info.feature_names[f], imp=float(imp[f]), lo=lo, hi=hi, v=round(med) if kind != "continuous" else med, step=step, kind=kind, nsplit=len(thr), thr=sorted(thr), **_marginal(info, A, f, thr))) samples = [] if info.X is not None and max_samples: rng = np.random.default_rng(0) rows = rng.choice(len(info.X), min(max_samples, len(info.X)), replace=False) for r in rows: truth = None if info.y is not None: truth = info.class_names[int(info.y[r])] if info.is_clf else fmt(info.y[r], max(precision, 4)) samples.append(dict(x={str(f): float(info.X[r, f]) for f in used}, y=truth, row=int(r))) sub = subtitle if sub is None: sub = info.summary() if info.is_clf: note = {"sum": "leaf values add up to the log-odds", "mean": "the trees' predictions are averaged"}.get(info.combine) legend = class_legend(info.class_names, P, note) else: label = "leaf value, added across trees" if info.combine == "sum" else f"predicted {info.target_name}" legend = [dict(k="grad", label=label, lo=info.value_range[0], hi=info.value_range[1], stops=[f"var(--dti-s{j})" for j in range(9)])] crit = info.criterion.replace("_", " ") or "impurity" if imp_kind == "impurity": imp_note = (f"<b>%</b> is the feature's share of the model's total {esc(crit)} reduction, summed over its " "splits and weighted by samples (sklearn's <code>feature_importances_</code> definition). ") else: imp_note = "<b>%</b> is the feature's share of all samples reaching a split on it (impurities were not available). " imp_note += "Bars are relative to the top feature; ticks under each distribution mark split thresholds." data = dict( nodes=nodes, feats=feats, samples=samples, task=info.task, classes=[dict(name=n, col=P.cls(k)) for k, n in enumerate(info.class_names or [])], legend=legend, impNote=imp_note, fileName=info.model_name, target=info.target_name, orientation=resolve_orientation(orientation, info), initialDepth=initial_depth if initial_depth is not None else (99 if len(nodes) <= 63 else 3), maxDepth=info.max_depth, simple=bool(info.n_leaves > 32 if simple == "auto" else simple), defs=defs(uid, P), sig=precision, bandMax=BAND, roots=info.roots, shown=list(info.shown_roots or info.roots), combine=info.combine, layout=info.layout, link=info.link, intercept=info.intercept, ) page_title = title or f"{info.model_name} ({info.n_leaves} leaves)" out = build_page("tree", title=page_title, subtitle=sub, theme=theme, data=data) return InteractiveTree(out, height) def text(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, precision=4, max_trees=3)-
Write a fitted model out as readable text: what
print(model)shows for imodels estimators.Trees print their conditions and leaf outcomes, rule lists their IF / ELSE IF rows, rule sets their rules ranked by effect, scorecards their points and the risk of each total, and additive models a sparkline of each shape function.
Parameters
model:estimator- A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features.
X:array-likeofshape (n_samples, n_features), optional- Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally.
y:array-likeofshape (n_samples,), optional- Targets for
X, used for class mixes and target ranges. feature_names:listofstr, optional- Feature names. Default: the columns of
X, else the names the model was fitted with. class_names:listordict, optional- Class names, as a list or as a dict from class label to name. Default: the model's classes.
target_name:str, optional- Name of the target (regression). Default: the name of
y, else "target". precision:int, default=4- Significant digits for values. Thresholds get more digits where two would print alike.
max_trees:int, default=3- For ensembles, how many trees to write out.
Returns
str- The model as text.
Expand source code
def text(model, X=None, y=None, *, feature_names=None, class_names=None, target_name=None, precision=4, max_trees=3): """Write a fitted model out as readable text: what ``print(model)`` shows for imodels estimators. Trees print their conditions and leaf outcomes, rule lists their IF / ELSE IF rows, rule sets their rules ranked by effect, scorecards their points and the risk of each total, and additive models a sparkline of each shape function. Parameters ---------- model : estimator A fitted model: an imodels estimator (tree, sum of trees, rule list, rule set, scoring system or additive model) or a scikit-learn tree, forest, gradient-boosting, linear or isotonic model, or a Pipeline whose earlier steps only scale features. X : array-like of shape (n_samples, n_features), optional Training (or held-out) data. With it, split nodes show their feature's distribution, rules show their coverage, and thresholds on binary or integer features read naturally. y : array-like of shape (n_samples,), optional Targets for ``X``, used for class mixes and target ranges. feature_names : list of str, optional Feature names. Default: the columns of ``X``, else the names the model was fitted with. class_names : list or dict, optional Class names, as a list or as a dict from class label to name. Default: the model's classes. target_name : str, optional Name of the target (regression). Default: the name of ``y``, else "target". precision : int, default=4 Significant digits for values. Thresholds get more digits where two would print alike. max_trees : int, default=3 For ensembles, how many trees to write out. Returns ------- str The model as text. """ from ._adapt import to_view view = to_view(model, X, y, feature_names, class_names, target_name) extra = None if hasattr(model, "optimal_") and hasattr(model, "objective_"): # FastSmallTree status = "certified optimal" if model.optimal_ else "not certified optimal (the time limit was reached)" extra = f"{status}, objective {model.objective_:.4g}" return model_text(view, precision, max_trees, extra)
Classes
class InteractiveTree (html_text, height=720)-
An interactive page for a model, returned by
interactive(): one self-contained HTML file. Displays inline in Jupyter.Attributes
html:str- The page as HTML text.
height:int- Height in pixels when shown inline in Jupyter.
Expand source code
class InteractiveTree: """An interactive page for a model, returned by `interactive`: one self-contained HTML file. Displays inline in Jupyter. Attributes ---------- html : str The page as HTML text. height : int Height in pixels when shown inline in Jupyter. """ def __init__(self, html_text, height=720): self.html = html_text self.height = height def save(self, path): """Write the page to an .html file. Parameters ---------- path : str or path-like Output file, ending in .html or .htm. Returns ------- path : str or path-like The path written. """ if os.path.splitext(str(path))[1].lower() not in (".html", ".htm"): raise ValueError("Interactive trees save to .html") with open(path, "w", encoding="utf-8") as f: f.write(self.html) return path def _repr_html_(self): return (f'<iframe srcdoc="{html.escape(self.html, quote=True)}" ' f'style="width:100%;height:{self.height}px;border:0;border-radius:12px" loading="lazy"></iframe>')Methods
def save(self, path)-
Write the page to an .html file.
Parameters
path:strorpath-like- Output file, ending in .html or .htm.
Returns
path:strorpath-like- The path written.
Expand source code
def save(self, path): """Write the page to an .html file. Parameters ---------- path : str or path-like Output file, ending in .html or .htm. Returns ------- path : str or path-like The path written. """ if os.path.splitext(str(path))[1].lower() not in (".html", ".htm"): raise ValueError("Interactive trees save to .html") with open(path, "w", encoding="utf-8") as f: f.write(self.html) return path
class TreeFigure (svg)-
A static figure of a model, returned by
draw(). Displays inline in Jupyter.Attributes
svg:str- The figure as SVG text.
Expand source code
class TreeFigure: """A static figure of a model, returned by `draw`. Displays inline in Jupyter. Attributes ---------- svg : str The figure as SVG text. """ def __init__(self, svg): self.svg = svg def _repr_svg_(self): return self.svg def __str__(self): return self.svg def save(self, path, scale=2.0): """Write the figure to a file. Parameters ---------- path : str or path-like Output file; its extension picks the format: .svg, .png, .pdf or .html. PNG and PDF need cairosvg. scale : float, default=2.0 Resolution multiplier for PNG (PDF is vector, so it is unaffected). Returns ------- path : str or path-like The path written. """ ext = os.path.splitext(str(path))[1].lower() if ext == ".svg": with open(path, "w", encoding="utf-8") as f: f.write(self.svg) elif ext in (".png", ".pdf"): try: import cairosvg except ImportError as e: # pragma: no cover raise ImportError("PNG/PDF export needs cairosvg: pip install cairosvg") from e fn = cairosvg.svg2png if ext == ".png" else cairosvg.svg2pdf from ._text import EXPORT_FONT, FONT_STACK, esc svg = self.svg.replace(f'font-family="{esc(FONT_STACK)}"', f'font-family="{EXPORT_FONT}"', 1) fn(bytestring=svg.encode("utf-8"), write_to=str(path), scale=scale if ext == ".png" else 1.0) elif ext in (".html", ".htm"): with open(path, "w", encoding="utf-8") as f: f.write("<!doctype html><meta charset='utf-8'><title>Model</title>" "<body style='margin:0;display:flex;justify-content:center'>" + self.svg + "</body>") else: raise ValueError(f"Unsupported extension {ext!r}; use .svg, .png, .pdf or .html") return pathMethods
def save(self, path, scale=2.0)-
Write the figure to a file.
Parameters
path:strorpath-like- Output file; its extension picks the format: .svg, .png, .pdf or .html. PNG and PDF need cairosvg.
scale:float, default=2.0- Resolution multiplier for PNG (PDF is vector, so it is unaffected).
Returns
path:strorpath-like- The path written.
Expand source code
def save(self, path, scale=2.0): """Write the figure to a file. Parameters ---------- path : str or path-like Output file; its extension picks the format: .svg, .png, .pdf or .html. PNG and PDF need cairosvg. scale : float, default=2.0 Resolution multiplier for PNG (PDF is vector, so it is unaffected). Returns ------- path : str or path-like The path written. """ ext = os.path.splitext(str(path))[1].lower() if ext == ".svg": with open(path, "w", encoding="utf-8") as f: f.write(self.svg) elif ext in (".png", ".pdf"): try: import cairosvg except ImportError as e: # pragma: no cover raise ImportError("PNG/PDF export needs cairosvg: pip install cairosvg") from e fn = cairosvg.svg2png if ext == ".png" else cairosvg.svg2pdf from ._text import EXPORT_FONT, FONT_STACK, esc svg = self.svg.replace(f'font-family="{esc(FONT_STACK)}"', f'font-family="{EXPORT_FONT}"', 1) fn(bytestring=svg.encode("utf-8"), write_to=str(path), scale=scale if ext == ".png" else 1.0) elif ext in (".html", ".htm"): with open(path, "w", encoding="utf-8") as f: f.write("<!doctype html><meta charset='utf-8'><title>Model</title>" "<body style='margin:0;display:flex;justify-content:center'>" + self.svg + "</body>") else: raise ValueError(f"Unsupported extension {ext!r}; use .svg, .png, .pdf or .html") return path