figkit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- figkit/__init__.py +114 -0
- figkit/audit.py +611 -0
- figkit/colors.py +304 -0
- figkit/component.py +110 -0
- figkit/components.py +727 -0
- figkit/connectors.py +689 -0
- figkit/core.py +1021 -0
- figkit/export.py +301 -0
- figkit/figure.py +312 -0
- figkit/fonts.py +462 -0
- figkit/frame.py +376 -0
- figkit/geom.py +497 -0
- figkit/image.py +308 -0
- figkit/layout.py +529 -0
- figkit/mathtext.py +320 -0
- figkit/paint.py +136 -0
- figkit/py.typed +0 -0
- figkit/shapes.py +764 -0
- figkit/style.py +502 -0
- figkit/svgdoc.py +133 -0
- figkit/svgpath.py +470 -0
- figkit/text.py +625 -0
- figkit/themes.py +163 -0
- figkit-0.1.0.dist-info/METADATA +277 -0
- figkit-0.1.0.dist-info/RECORD +28 -0
- figkit-0.1.0.dist-info/WHEEL +5 -0
- figkit-0.1.0.dist-info/licenses/LICENSE +21 -0
- figkit-0.1.0.dist-info/top_level.txt +1 -0
figkit/components.py
ADDED
|
@@ -0,0 +1,727 @@
|
|
|
1
|
+
"""Higher-level building blocks: panels, matrices, braces, legends, tables."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
|
|
7
|
+
from .colors import colormap, contrast_color, to_hex
|
|
8
|
+
from .component import Component
|
|
9
|
+
from .core import Element, Group
|
|
10
|
+
from .style import Style
|
|
11
|
+
from .geom import BBox, Point, _expand_spec, to_point
|
|
12
|
+
from .paint import paint_attrs
|
|
13
|
+
from .shapes import Box, Marker
|
|
14
|
+
from .svgdoc import Node, RenderContext
|
|
15
|
+
from .svgpath import fmt
|
|
16
|
+
from .text import Text
|
|
17
|
+
|
|
18
|
+
_GROUP_KEYS = frozenset({"name", "z", "visible", "opacity", "transform",
|
|
19
|
+
"clip", "add", "theme", "style"})
|
|
20
|
+
|
|
21
|
+
__all__ = ["Panel", "Matrix", "Vector", "Heatmap", "Brace", "Bracket",
|
|
22
|
+
"Legend", "Table", "Callout", "Spacer", "ColorBar",
|
|
23
|
+
"LabelledMatrix"]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
# ==========================================================================
|
|
27
|
+
# Panel — a box that tracks a set of elements
|
|
28
|
+
# ==========================================================================
|
|
29
|
+
|
|
30
|
+
class Panel(Box):
|
|
31
|
+
"""A rectangle sized around other elements, drawn behind them.
|
|
32
|
+
|
|
33
|
+
Created by :func:`figkit.layout.fit`, but usable directly:
|
|
34
|
+
|
|
35
|
+
>>> Panel([box_a, box_b], pad=16, label="Fused operation", dash=True)
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
role = "panel"
|
|
39
|
+
default_padding = 0
|
|
40
|
+
|
|
41
|
+
def __init__(self, targets=None, pad=16, label=None,
|
|
42
|
+
label_pos: str = "above_left",
|
|
43
|
+
label_gap: float = 6.0, label_style=None, follow: bool = True,
|
|
44
|
+
**kw):
|
|
45
|
+
if targets is None:
|
|
46
|
+
targets = []
|
|
47
|
+
elif isinstance(targets, Element):
|
|
48
|
+
targets = [targets]
|
|
49
|
+
self.targets = list(targets)
|
|
50
|
+
self.pad = _expand_spec(pad)
|
|
51
|
+
self.follow = follow
|
|
52
|
+
self.label_pos = str(label_pos).lower()
|
|
53
|
+
self.label_gap = float(label_gap)
|
|
54
|
+
kw.setdefault("z", -1000)
|
|
55
|
+
kw.pop("text", None)
|
|
56
|
+
super().__init__(None, **kw)
|
|
57
|
+
self._panel_label: Text | None = None
|
|
58
|
+
if label:
|
|
59
|
+
self._panel_label = Text(label, add=False, style=label_style)
|
|
60
|
+
self._panel_label.parent = self
|
|
61
|
+
self._panel_label.role = "label"
|
|
62
|
+
self.invalidate()
|
|
63
|
+
|
|
64
|
+
def add_target(self, *elements) -> "Panel":
|
|
65
|
+
for el in elements:
|
|
66
|
+
if el is not None and el is not self:
|
|
67
|
+
self.targets.append(el)
|
|
68
|
+
return self.invalidate()
|
|
69
|
+
|
|
70
|
+
def _measure(self) -> None:
|
|
71
|
+
if self.follow and self.targets:
|
|
72
|
+
self._dirty = False
|
|
73
|
+
bb = BBox.union_all([t.ink_bbox for t in self.targets
|
|
74
|
+
if t.visible and t is not self])
|
|
75
|
+
if bb is not None:
|
|
76
|
+
t, r, b, l = self.pad
|
|
77
|
+
bb = bb.expand(0, top=t, right=r, bottom=b, left=l)
|
|
78
|
+
inv = self.world_matrix()
|
|
79
|
+
if not inv.is_identity:
|
|
80
|
+
bb = inv.inverse().apply_bbox(bb)
|
|
81
|
+
self._x, self._y, self._w, self._h = bb.x, bb.y, bb.w, bb.h
|
|
82
|
+
else:
|
|
83
|
+
super()._measure()
|
|
84
|
+
if self._panel_label is not None:
|
|
85
|
+
self._place_panel_label()
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def local_bbox(self) -> BBox:
|
|
89
|
+
if self.follow and self.targets:
|
|
90
|
+
self._dirty = True # the elements we track may have moved
|
|
91
|
+
self._ensure()
|
|
92
|
+
return BBox(self._x, self._y, self._w or 0.0, self._h or 0.0)
|
|
93
|
+
|
|
94
|
+
def _place_panel_label(self) -> None:
|
|
95
|
+
lbl = self._panel_label
|
|
96
|
+
bb = BBox(self._x, self._y, self._w or 0, self._h or 0)
|
|
97
|
+
g = self.label_gap
|
|
98
|
+
pos = self.label_pos
|
|
99
|
+
inner = {"nw": ("nw", g, g), "n": ("n", 0, g), "ne": ("ne", -g, g),
|
|
100
|
+
"sw": ("sw", g, -g), "s": ("s", 0, -g), "se": ("se", -g, -g),
|
|
101
|
+
"w": ("w", g, 0), "e": ("e", -g, 0), "center": ("center", 0, 0)}
|
|
102
|
+
if pos in inner:
|
|
103
|
+
anchor, dx, dy = inner[pos]
|
|
104
|
+
p = bb.anchor(anchor)
|
|
105
|
+
lbl.place_local(p.x + dx, p.y + dy, anchor=anchor)
|
|
106
|
+
return
|
|
107
|
+
outer = {"above": ("s", "n", 0, -g), "top": ("s", "n", 0, -g),
|
|
108
|
+
"below": ("n", "s", 0, g), "bottom": ("n", "s", 0, g),
|
|
109
|
+
"left": ("e", "w", -g, 0), "right": ("w", "e", g, 0),
|
|
110
|
+
"above_left": ("sw", "nw", 0, -g),
|
|
111
|
+
"above_right": ("se", "ne", 0, -g)}
|
|
112
|
+
anchor, side, dx, dy = outer.get(pos, ("s", "n", 0, -g))
|
|
113
|
+
p = bb.anchor(side)
|
|
114
|
+
lbl.place_local(p.x + dx, p.y + dy, anchor=anchor)
|
|
115
|
+
|
|
116
|
+
@property
|
|
117
|
+
def label(self) -> Text | None:
|
|
118
|
+
self._ensure()
|
|
119
|
+
return self._panel_label
|
|
120
|
+
|
|
121
|
+
def _render_content(self, ctx: RenderContext):
|
|
122
|
+
self._ensure()
|
|
123
|
+
nodes = list(self.shape_nodes(ctx, self.local_bbox) or [])
|
|
124
|
+
if self._panel_label is not None:
|
|
125
|
+
n = self._panel_label.render(ctx)
|
|
126
|
+
if n is not None:
|
|
127
|
+
nodes.append(n)
|
|
128
|
+
return nodes or None
|
|
129
|
+
|
|
130
|
+
@property
|
|
131
|
+
def ink_bbox(self) -> BBox:
|
|
132
|
+
bb = super().ink_bbox
|
|
133
|
+
if self._panel_label is not None:
|
|
134
|
+
bb = bb.union(self._panel_label.ink_bbox)
|
|
135
|
+
return bb
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
# ==========================================================================
|
|
139
|
+
# Matrix / heatmap
|
|
140
|
+
# ==========================================================================
|
|
141
|
+
|
|
142
|
+
class Matrix(Group):
|
|
143
|
+
"""A grid of coloured cells — feature vectors, attention maps, matrices.
|
|
144
|
+
|
|
145
|
+
>>> Matrix([[0.1, 0.9], [0.5, 0.2]], cell=18, cmap="viridis")
|
|
146
|
+
>>> Matrix(colors=[["#eee", "#333"]], cell=14)
|
|
147
|
+
|
|
148
|
+
``m.cell(i, j)`` returns the cell as a real element, so you can point an
|
|
149
|
+
arrow at it or restyle it.
|
|
150
|
+
"""
|
|
151
|
+
|
|
152
|
+
role = "matrix"
|
|
153
|
+
|
|
154
|
+
def __init__(self, values=None, *, colors=None, rows: int = None,
|
|
155
|
+
cols: int = None, cell=16, gap: float = 0.0,
|
|
156
|
+
cmap="viridis", vmin: float = None, vmax: float = None,
|
|
157
|
+
radius: float = None, show_values: bool = False,
|
|
158
|
+
value_fmt="{:.2f}", value_color=None, value_size: float = None,
|
|
159
|
+
border=None, cell_style=None, x: float = 0.0, y: float = 0.0,
|
|
160
|
+
**kw):
|
|
161
|
+
self._grid_values = _as_grid(values, rows, cols)
|
|
162
|
+
self._grid_colors = _as_grid(colors, rows, cols) if colors else None
|
|
163
|
+
if self._grid_values is None and self._grid_colors is None:
|
|
164
|
+
self._grid_values = [[0.0] * (cols or 1) for _ in range(rows or 1)]
|
|
165
|
+
src = self._grid_colors or self._grid_values
|
|
166
|
+
self.n_rows = len(src)
|
|
167
|
+
self.n_cols = max(len(r) for r in src) if src else 0
|
|
168
|
+
cw, ch = (cell if isinstance(cell, (tuple, list)) else (cell, cell))
|
|
169
|
+
self.cell_w, self.cell_h = float(cw), float(ch)
|
|
170
|
+
self.gap = float(gap)
|
|
171
|
+
self.cmap = cmap
|
|
172
|
+
self._cells: list = []
|
|
173
|
+
self._value_labels: list = []
|
|
174
|
+
# Paint properties passed to Matrix(...) style the *cells*, which is
|
|
175
|
+
# what people mean by `Matrix(vals, stroke="#333")`.
|
|
176
|
+
group_kw = {k: v for k, v in kw.items() if k in _GROUP_KEYS}
|
|
177
|
+
cell_kw = {k: v for k, v in kw.items() if k not in _GROUP_KEYS}
|
|
178
|
+
cell_style = Style(cell_style, **cell_kw) if cell_kw else cell_style
|
|
179
|
+
super().__init__(**group_kw)
|
|
180
|
+
|
|
181
|
+
flat = [v for row in (self._grid_values or []) for v in row
|
|
182
|
+
if v is not None]
|
|
183
|
+
lo = min(flat) if (flat and vmin is None) else vmin
|
|
184
|
+
hi = max(flat) if (flat and vmax is None) else vmax
|
|
185
|
+
self.vmin, self.vmax = lo, hi
|
|
186
|
+
|
|
187
|
+
for i in range(self.n_rows):
|
|
188
|
+
row_cells = []
|
|
189
|
+
for j in range(self.n_cols):
|
|
190
|
+
cx = x + j * (self.cell_w + self.gap)
|
|
191
|
+
cy = y + i * (self.cell_h + self.gap)
|
|
192
|
+
fill = self._cell_color(i, j)
|
|
193
|
+
cell_el = Box(None, cx, cy, self.cell_w, self.cell_h,
|
|
194
|
+
style=cell_style, fill=fill, padding=0,
|
|
195
|
+
radius=radius if radius is not None else None,
|
|
196
|
+
add=False)
|
|
197
|
+
cell_el.role = "matrix"
|
|
198
|
+
self.add(cell_el)
|
|
199
|
+
row_cells.append(cell_el)
|
|
200
|
+
if show_values and self._grid_values:
|
|
201
|
+
v = self._value(i, j)
|
|
202
|
+
if v is None:
|
|
203
|
+
continue
|
|
204
|
+
txt = value_fmt.format(v) if not callable(value_fmt) \
|
|
205
|
+
else value_fmt(v)
|
|
206
|
+
col = value_color or contrast_color(fill)
|
|
207
|
+
lbl = Text(txt, add=False, color=col,
|
|
208
|
+
font_size=value_size or max(7.0, self.cell_h * 0.38))
|
|
209
|
+
lbl.parent = self
|
|
210
|
+
lbl.role = "label"
|
|
211
|
+
self.add(lbl)
|
|
212
|
+
lbl.center_at(cell_el.bbox.cx, cell_el.bbox.cy)
|
|
213
|
+
self._value_labels.append(lbl)
|
|
214
|
+
self._cells.append(row_cells)
|
|
215
|
+
|
|
216
|
+
if border:
|
|
217
|
+
bb = self.bbox
|
|
218
|
+
frame = Box(None, bb.x, bb.y, bb.w, bb.h, fill="none",
|
|
219
|
+
stroke=border if isinstance(border, str) else "#333",
|
|
220
|
+
padding=0, radius=0, z=10, add=False)
|
|
221
|
+
self.add(frame)
|
|
222
|
+
self.border = frame
|
|
223
|
+
|
|
224
|
+
# -- data ------------------------------------------------------------
|
|
225
|
+
def _value(self, i: int, j: int):
|
|
226
|
+
if not self._grid_values:
|
|
227
|
+
return None
|
|
228
|
+
row = self._grid_values[i] if i < len(self._grid_values) else []
|
|
229
|
+
return row[j] if j < len(row) else None
|
|
230
|
+
|
|
231
|
+
def _cell_color(self, i: int, j: int) -> str:
|
|
232
|
+
if self._grid_colors:
|
|
233
|
+
row = self._grid_colors[i] if i < len(self._grid_colors) else []
|
|
234
|
+
if j < len(row) and row[j] is not None:
|
|
235
|
+
return to_hex(row[j])
|
|
236
|
+
v = self._value(i, j)
|
|
237
|
+
if v is None:
|
|
238
|
+
return "#dddddd"
|
|
239
|
+
lo, hi = self.vmin, self.vmax
|
|
240
|
+
t = 0.5 if (lo is None or hi is None or hi == lo) else (v - lo) / (hi - lo)
|
|
241
|
+
return colormap(self.cmap, t)
|
|
242
|
+
|
|
243
|
+
# -- access ----------------------------------------------------------
|
|
244
|
+
def cell(self, i: int, j: int = 0) -> Box:
|
|
245
|
+
"""The cell element at row ``i``, column ``j`` (negatives wrap)."""
|
|
246
|
+
return self._cells[i][j]
|
|
247
|
+
|
|
248
|
+
@property
|
|
249
|
+
def cells(self) -> list:
|
|
250
|
+
return self._cells
|
|
251
|
+
|
|
252
|
+
def row(self, i: int) -> list:
|
|
253
|
+
return list(self._cells[i])
|
|
254
|
+
|
|
255
|
+
def col(self, j: int) -> list:
|
|
256
|
+
return [r[j] for r in self._cells]
|
|
257
|
+
|
|
258
|
+
def shape(self) -> tuple:
|
|
259
|
+
return (self.n_rows, self.n_cols)
|
|
260
|
+
|
|
261
|
+
def highlight(self, i: int, j: int = 0, **style) -> Box:
|
|
262
|
+
"""Restyle one cell (e.g. ``highlight(0, 2, stroke='red', lw=2)``)."""
|
|
263
|
+
c = self.cell(i, j)
|
|
264
|
+
c.restyle(**style)
|
|
265
|
+
c.to_front()
|
|
266
|
+
return c
|
|
267
|
+
|
|
268
|
+
|
|
269
|
+
def Vector(values=None, *, orient: str = "v", **kw) -> Matrix:
|
|
270
|
+
"""A 1-D :class:`Matrix` — the little feature-vector strips in ML figures."""
|
|
271
|
+
if values is None:
|
|
272
|
+
values = []
|
|
273
|
+
flat = list(values)
|
|
274
|
+
if str(orient).lower().startswith("v"):
|
|
275
|
+
grid = [[v] for v in flat]
|
|
276
|
+
else:
|
|
277
|
+
grid = [flat]
|
|
278
|
+
colors = kw.pop("colors", None)
|
|
279
|
+
if colors is not None:
|
|
280
|
+
colors = [[c] for c in colors] if str(orient).lower().startswith("v") \
|
|
281
|
+
else [list(colors)]
|
|
282
|
+
return Matrix(None, colors=colors, **kw)
|
|
283
|
+
return Matrix(grid, **kw)
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
Heatmap = Matrix
|
|
287
|
+
|
|
288
|
+
|
|
289
|
+
class ColorBar(Group):
|
|
290
|
+
"""A gradient strip with min/max labels, for heatmap legends."""
|
|
291
|
+
|
|
292
|
+
role = "group"
|
|
293
|
+
|
|
294
|
+
def __init__(self, cmap="viridis", vmin=0.0, vmax=1.0, *, w: float = 14.0,
|
|
295
|
+
h: float = 120.0, steps: int = 24, orient: str = "v",
|
|
296
|
+
labels=True, label_fmt="{:.2g}", x: float = 0.0,
|
|
297
|
+
y: float = 0.0, **kw):
|
|
298
|
+
super().__init__(**kw)
|
|
299
|
+
vertical = str(orient).lower().startswith("v")
|
|
300
|
+
n = max(2, int(steps))
|
|
301
|
+
for i in range(n):
|
|
302
|
+
t = i / (n - 1)
|
|
303
|
+
col = colormap(cmap, 1.0 - t if vertical else t)
|
|
304
|
+
if vertical:
|
|
305
|
+
seg = Box(None, x, y + t * h, w, h / n + 0.6, fill=col,
|
|
306
|
+
stroke="none", padding=0, radius=0, add=False)
|
|
307
|
+
else:
|
|
308
|
+
seg = Box(None, x + t * w, y, w / n + 0.6, h, fill=col,
|
|
309
|
+
stroke="none", padding=0, radius=0, add=False)
|
|
310
|
+
self.add(seg)
|
|
311
|
+
frame = Box(None, x, y, w, h, fill="none", stroke="#6b7280",
|
|
312
|
+
stroke_width=0.8, padding=0, radius=0, add=False)
|
|
313
|
+
self.add(frame)
|
|
314
|
+
if labels:
|
|
315
|
+
lo = Text(label_fmt.format(vmin), add=False, font_size=10)
|
|
316
|
+
hi = Text(label_fmt.format(vmax), add=False, font_size=10)
|
|
317
|
+
self.add(lo, hi)
|
|
318
|
+
if vertical:
|
|
319
|
+
hi.at(x + w + 5, y, anchor="w")
|
|
320
|
+
lo.at(x + w + 5, y + h, anchor="w")
|
|
321
|
+
else:
|
|
322
|
+
lo.at(x, y + h + 4, anchor="n")
|
|
323
|
+
hi.at(x + w, y + h + 4, anchor="n")
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
# ==========================================================================
|
|
327
|
+
# Braces & brackets
|
|
328
|
+
# ==========================================================================
|
|
329
|
+
|
|
330
|
+
class Brace(Element):
|
|
331
|
+
"""A curly brace spanning two points, optionally labelled.
|
|
332
|
+
|
|
333
|
+
>>> Brace(box_a.nw, box_c.ne, depth=12, label="encoder")
|
|
334
|
+
"""
|
|
335
|
+
|
|
336
|
+
role = "brace"
|
|
337
|
+
STROKE_WIDTH_ALIAS = True
|
|
338
|
+
|
|
339
|
+
def __init__(self, start, end, *, depth: float = 12.0, side: str = "auto",
|
|
340
|
+
label=None, label_gap: float = 6.0, label_style=None,
|
|
341
|
+
sharpness: float = 0.55, **kw):
|
|
342
|
+
self.start_ref = start
|
|
343
|
+
self.end_ref = end
|
|
344
|
+
self.depth = float(depth)
|
|
345
|
+
self.side = str(side).lower()
|
|
346
|
+
self.sharpness = float(sharpness)
|
|
347
|
+
self.label_gap = float(label_gap)
|
|
348
|
+
self._label: Text | None = None
|
|
349
|
+
super().__init__(0, 0, None, None, **kw)
|
|
350
|
+
if label:
|
|
351
|
+
self._label = Text(label, add=False, style=label_style)
|
|
352
|
+
self._label.parent = self
|
|
353
|
+
self._label.role = "label"
|
|
354
|
+
|
|
355
|
+
def _normal(self, p0: Point, p1: Point) -> Point:
|
|
356
|
+
v = (p1 - p0).normalized()
|
|
357
|
+
n = Point(-v.y, v.x)
|
|
358
|
+
if self.side in ("auto", ""):
|
|
359
|
+
return n
|
|
360
|
+
wanted = {"up": Point(0, -1), "above": Point(0, -1),
|
|
361
|
+
"down": Point(0, 1), "below": Point(0, 1),
|
|
362
|
+
"left": Point(-1, 0), "right": Point(1, 0)}.get(self.side)
|
|
363
|
+
if wanted is None:
|
|
364
|
+
return n
|
|
365
|
+
return n if n.dot(wanted) >= 0 else -n
|
|
366
|
+
|
|
367
|
+
def path_data(self, bb: BBox = None) -> str:
|
|
368
|
+
p0 = to_point(self.start_ref)
|
|
369
|
+
p1 = to_point(self.end_ref)
|
|
370
|
+
n = self._normal(p0, p1)
|
|
371
|
+
d = self.depth
|
|
372
|
+
mid = p0.lerp(p1, 0.5)
|
|
373
|
+
tip = mid + n * d
|
|
374
|
+
q = self.sharpness
|
|
375
|
+
a1 = p0.lerp(mid, q) + n * d * 0.5
|
|
376
|
+
a2 = mid.lerp(p0, 1 - q) + n * d * 0.5
|
|
377
|
+
b1 = mid.lerp(p1, 1 - q) + n * d * 0.5
|
|
378
|
+
b2 = p1.lerp(mid, q) + n * d * 0.5
|
|
379
|
+
h0 = p0 + n * d * 0.5
|
|
380
|
+
h1 = p1 + n * d * 0.5
|
|
381
|
+
return (f"M{fmt(p0.x)} {fmt(p0.y)}"
|
|
382
|
+
f"Q{fmt(h0.x)} {fmt(h0.y)} {fmt(a1.x)} {fmt(a1.y)}"
|
|
383
|
+
f"L{fmt(a2.x)} {fmt(a2.y)}"
|
|
384
|
+
f"Q{fmt(tip.x)} {fmt(tip.y)} {fmt(b1.x)} {fmt(b1.y)}"
|
|
385
|
+
f"L{fmt(b2.x)} {fmt(b2.y)}"
|
|
386
|
+
f"Q{fmt(h1.x)} {fmt(h1.y)} {fmt(p1.x)} {fmt(p1.y)}")
|
|
387
|
+
|
|
388
|
+
@property
|
|
389
|
+
def tip(self) -> Point:
|
|
390
|
+
p0 = to_point(self.start_ref)
|
|
391
|
+
p1 = to_point(self.end_ref)
|
|
392
|
+
return p0.lerp(p1, 0.5) + self._normal(p0, p1) * self.depth
|
|
393
|
+
|
|
394
|
+
def _measure(self) -> None:
|
|
395
|
+
from .svgpath import path_bbox
|
|
396
|
+
x0, y0, x1, y1 = path_bbox(self.path_data())
|
|
397
|
+
self._x, self._y, self._w, self._h = x0, y0, x1 - x0, y1 - y0
|
|
398
|
+
if self._label is not None:
|
|
399
|
+
p0 = to_point(self.start_ref)
|
|
400
|
+
p1 = to_point(self.end_ref)
|
|
401
|
+
n = self._normal(p0, p1)
|
|
402
|
+
target = self.tip + n * self.label_gap
|
|
403
|
+
anchor = "center"
|
|
404
|
+
if abs(n.x) > abs(n.y):
|
|
405
|
+
anchor = "w" if n.x > 0 else "e"
|
|
406
|
+
else:
|
|
407
|
+
anchor = "n" if n.y > 0 else "s"
|
|
408
|
+
self._label.at(target.x, target.y, anchor=anchor)
|
|
409
|
+
|
|
410
|
+
@property
|
|
411
|
+
def local_bbox(self) -> BBox:
|
|
412
|
+
self._dirty = True
|
|
413
|
+
self._ensure()
|
|
414
|
+
bb = BBox(self._x, self._y, self._w or 0.0, self._h or 0.0)
|
|
415
|
+
if self._label is not None:
|
|
416
|
+
bb = bb.union(self._label.local_bbox)
|
|
417
|
+
return bb
|
|
418
|
+
|
|
419
|
+
@property
|
|
420
|
+
def label(self) -> Text | None:
|
|
421
|
+
self._ensure()
|
|
422
|
+
return self._label
|
|
423
|
+
|
|
424
|
+
def _render_content(self, ctx: RenderContext):
|
|
425
|
+
attrs = paint_attrs(self, ctx)
|
|
426
|
+
attrs["fill"] = "none"
|
|
427
|
+
attrs.setdefault("stroke-linecap", "round")
|
|
428
|
+
nodes = [Node("path", d=self.path_data(), **attrs)]
|
|
429
|
+
if self._label is not None:
|
|
430
|
+
self._ensure()
|
|
431
|
+
n = self._label.render(ctx)
|
|
432
|
+
if n is not None:
|
|
433
|
+
nodes.append(n)
|
|
434
|
+
return nodes
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
class Bracket(Brace):
|
|
438
|
+
"""A square bracket instead of a curly one."""
|
|
439
|
+
|
|
440
|
+
def path_data(self, bb: BBox = None) -> str:
|
|
441
|
+
p0 = to_point(self.start_ref)
|
|
442
|
+
p1 = to_point(self.end_ref)
|
|
443
|
+
n = self._normal(p0, p1) * self.depth
|
|
444
|
+
return (f"M{fmt(p0.x)} {fmt(p0.y)}"
|
|
445
|
+
f"L{fmt((p0 + n).x)} {fmt((p0 + n).y)}"
|
|
446
|
+
f"L{fmt((p1 + n).x)} {fmt((p1 + n).y)}"
|
|
447
|
+
f"L{fmt(p1.x)} {fmt(p1.y)}")
|
|
448
|
+
|
|
449
|
+
@property
|
|
450
|
+
def tip(self) -> Point:
|
|
451
|
+
p0 = to_point(self.start_ref)
|
|
452
|
+
p1 = to_point(self.end_ref)
|
|
453
|
+
return p0.lerp(p1, 0.5) + self._normal(p0, p1) * self.depth
|
|
454
|
+
|
|
455
|
+
|
|
456
|
+
# ==========================================================================
|
|
457
|
+
# Legend & table
|
|
458
|
+
# ==========================================================================
|
|
459
|
+
|
|
460
|
+
class Legend(Group):
|
|
461
|
+
"""A colour/marker legend.
|
|
462
|
+
|
|
463
|
+
>>> Legend([("train", "#4C72B0"), ("val", "#DD8452")], marker="square")
|
|
464
|
+
"""
|
|
465
|
+
|
|
466
|
+
role = "group"
|
|
467
|
+
|
|
468
|
+
def __init__(self, entries, *, marker: str = "square", swatch: float = 12.0,
|
|
469
|
+
gap: float = 7.0, row_gap: float = 6.0, cols: int = 1,
|
|
470
|
+
col_gap: float = 22.0, font_size: float = None,
|
|
471
|
+
x: float = 0.0, y: float = 0.0, **kw):
|
|
472
|
+
super().__init__(**kw)
|
|
473
|
+
self.rows: list = []
|
|
474
|
+
items = []
|
|
475
|
+
for entry in entries:
|
|
476
|
+
if isinstance(entry, (tuple, list)):
|
|
477
|
+
label, color = entry[0], entry[1]
|
|
478
|
+
opts = entry[2] if len(entry) > 2 else {}
|
|
479
|
+
else:
|
|
480
|
+
label, color, opts = str(entry), "#888888", {}
|
|
481
|
+
items.append((label, color, dict(opts)))
|
|
482
|
+
|
|
483
|
+
n = len(items)
|
|
484
|
+
per_col = math.ceil(n / max(1, cols))
|
|
485
|
+
col_x = x
|
|
486
|
+
widths = []
|
|
487
|
+
for c in range(cols):
|
|
488
|
+
chunk = items[c * per_col:(c + 1) * per_col]
|
|
489
|
+
cy = y
|
|
490
|
+
col_w = 0.0
|
|
491
|
+
for label, color, opts in chunk:
|
|
492
|
+
shape = opts.pop("marker", marker)
|
|
493
|
+
sw = float(opts.pop("swatch", swatch))
|
|
494
|
+
if shape in ("square", "rect", "box"):
|
|
495
|
+
sym = Box(None, col_x, cy, sw, sw, fill=color,
|
|
496
|
+
stroke=opts.pop("stroke", "none"), padding=0,
|
|
497
|
+
radius=opts.pop("radius", 2), add=False)
|
|
498
|
+
elif shape in ("line", "dash"):
|
|
499
|
+
from .shapes import Line
|
|
500
|
+
sym = Line((col_x, cy + sw / 2), (col_x + sw * 1.4, cy + sw / 2),
|
|
501
|
+
stroke=color, stroke_width=opts.pop("lw", 2.4),
|
|
502
|
+
stroke_dash=opts.pop("dash", None), add=False)
|
|
503
|
+
else:
|
|
504
|
+
sym = Marker((col_x + sw / 2, cy + sw / 2), sw, shape,
|
|
505
|
+
fill=color, add=False)
|
|
506
|
+
txt = Text(label, add=False, align="left",
|
|
507
|
+
font_size=font_size, **opts)
|
|
508
|
+
txt.parent = self
|
|
509
|
+
self.add(sym, txt)
|
|
510
|
+
txt.at(sym.bbox.x1 + gap, sym.bbox.cy, anchor="w")
|
|
511
|
+
row_h = max(sym.bbox.h, txt.bbox.h)
|
|
512
|
+
col_w = max(col_w, txt.bbox.x1 - col_x)
|
|
513
|
+
self.rows.append((sym, txt))
|
|
514
|
+
cy += row_h + row_gap
|
|
515
|
+
widths.append(col_w)
|
|
516
|
+
col_x += col_w + col_gap
|
|
517
|
+
|
|
518
|
+
|
|
519
|
+
class Table(Group):
|
|
520
|
+
"""A simple table of text cells with optional header styling.
|
|
521
|
+
|
|
522
|
+
>>> Table([["model", "acc"], ["ours", "92.1"]], header=True)
|
|
523
|
+
"""
|
|
524
|
+
|
|
525
|
+
role = "group"
|
|
526
|
+
|
|
527
|
+
def __init__(self, rows, *, header: bool = True, col_widths=None,
|
|
528
|
+
row_height: float = None, padding=(6, 10), align="left",
|
|
529
|
+
header_style=None, cell_style=None, stripe=None,
|
|
530
|
+
grid_lines: bool = True, x: float = 0.0, y: float = 0.0, **kw):
|
|
531
|
+
super().__init__(**kw)
|
|
532
|
+
data = [[("" if c is None else str(c)) for c in row] for row in rows]
|
|
533
|
+
n_cols = max((len(r) for r in data), default=0)
|
|
534
|
+
data = [r + [""] * (n_cols - len(r)) for r in data]
|
|
535
|
+
pt, pr, pb, pl = _expand_spec(padding)
|
|
536
|
+
aligns = align if isinstance(align, (list, tuple)) else [align] * n_cols
|
|
537
|
+
aligns = list(aligns) + [aligns[-1] if aligns else "left"] * n_cols
|
|
538
|
+
|
|
539
|
+
probes = []
|
|
540
|
+
for i, row in enumerate(data):
|
|
541
|
+
probe_row = []
|
|
542
|
+
for j, cell in enumerate(row):
|
|
543
|
+
style = header_style if (header and i == 0) else cell_style
|
|
544
|
+
t = Text(cell, add=False, align=aligns[j], style=style)
|
|
545
|
+
if header and i == 0 and (style is None or "font_weight" not in
|
|
546
|
+
(style or {})):
|
|
547
|
+
t.restyle(bold=True)
|
|
548
|
+
t.parent = self
|
|
549
|
+
probe_row.append(t)
|
|
550
|
+
probes.append(probe_row)
|
|
551
|
+
|
|
552
|
+
if col_widths is None:
|
|
553
|
+
col_widths = [max((probes[i][j].bbox.w for i in range(len(data))),
|
|
554
|
+
default=0) + pl + pr for j in range(n_cols)]
|
|
555
|
+
else:
|
|
556
|
+
col_widths = list(col_widths)
|
|
557
|
+
rh = row_height or (max((t.bbox.h for row in probes for t in row),
|
|
558
|
+
default=14) + pt + pb)
|
|
559
|
+
|
|
560
|
+
self.cells: list = []
|
|
561
|
+
cy = y
|
|
562
|
+
for i, row in enumerate(probes):
|
|
563
|
+
cx = x
|
|
564
|
+
row_cells = []
|
|
565
|
+
for j, t in enumerate(row):
|
|
566
|
+
w = col_widths[j]
|
|
567
|
+
bg = None
|
|
568
|
+
if header and i == 0:
|
|
569
|
+
bg = "#eef0f3"
|
|
570
|
+
elif stripe and i % 2 == 0:
|
|
571
|
+
bg = stripe if isinstance(stripe, str) else "#f7f8fa"
|
|
572
|
+
cell = Box(None, cx, cy, w, rh, fill=bg or "none",
|
|
573
|
+
stroke="#d7dbe0" if grid_lines else "none",
|
|
574
|
+
stroke_width=0.8, radius=0, padding=0, add=False)
|
|
575
|
+
self.add(cell)
|
|
576
|
+
self.add(t)
|
|
577
|
+
a = aligns[j]
|
|
578
|
+
if a in ("left", "start"):
|
|
579
|
+
t.at(cx + pl, cy + rh / 2, anchor="w")
|
|
580
|
+
elif a in ("right", "end"):
|
|
581
|
+
t.at(cx + w - pr, cy + rh / 2, anchor="e")
|
|
582
|
+
else:
|
|
583
|
+
t.at(cx + w / 2, cy + rh / 2, anchor="center")
|
|
584
|
+
row_cells.append((cell, t))
|
|
585
|
+
cx += w
|
|
586
|
+
self.cells.append(row_cells)
|
|
587
|
+
cy += rh
|
|
588
|
+
|
|
589
|
+
def cell(self, i: int, j: int) -> Box:
|
|
590
|
+
return self.cells[i][j][0]
|
|
591
|
+
|
|
592
|
+
def cell_text(self, i: int, j: int) -> Text:
|
|
593
|
+
return self.cells[i][j][1]
|
|
594
|
+
|
|
595
|
+
|
|
596
|
+
class Callout(Box):
|
|
597
|
+
"""A rounded box with a pointer tail aimed at a target point."""
|
|
598
|
+
|
|
599
|
+
role = "box"
|
|
600
|
+
|
|
601
|
+
def __init__(self, text=None, target=None, *, tail: float = 12.0, **kw):
|
|
602
|
+
self.target_ref = target
|
|
603
|
+
self.tail_w = float(tail)
|
|
604
|
+
kw.setdefault("radius", 8)
|
|
605
|
+
super().__init__(text, **kw)
|
|
606
|
+
|
|
607
|
+
def shape_nodes(self, ctx: RenderContext, bb: BBox) -> list:
|
|
608
|
+
nodes = super().shape_nodes(ctx, bb)
|
|
609
|
+
if self.target_ref is None:
|
|
610
|
+
return nodes
|
|
611
|
+
tgt = to_point(self.target_ref)
|
|
612
|
+
n = Point(tgt.x - bb.cx, tgt.y - bb.cy)
|
|
613
|
+
if n.length == 0:
|
|
614
|
+
return nodes
|
|
615
|
+
base = bb.at_angle(math.degrees(math.atan2(n.y, n.x)))
|
|
616
|
+
perp = Point(-n.y, n.x).normalized() * (self.tail_w / 2.0)
|
|
617
|
+
attrs = paint_attrs(self, ctx, bbox=bb)
|
|
618
|
+
attrs["stroke"] = "none"
|
|
619
|
+
tail = Node("path", d=(f"M{fmt((base + perp).x)} {fmt((base + perp).y)}"
|
|
620
|
+
f"L{fmt(tgt.x)} {fmt(tgt.y)}"
|
|
621
|
+
f"L{fmt((base - perp).x)} {fmt((base - perp).y)}Z"),
|
|
622
|
+
**attrs)
|
|
623
|
+
return list(nodes) + [tail]
|
|
624
|
+
|
|
625
|
+
|
|
626
|
+
class Spacer(Element):
|
|
627
|
+
"""An invisible element that just occupies space in a stack."""
|
|
628
|
+
|
|
629
|
+
role = "group"
|
|
630
|
+
|
|
631
|
+
def __init__(self, w: float = 0.0, h: float = 0.0, **kw):
|
|
632
|
+
kw.setdefault("visible", False)
|
|
633
|
+
super().__init__(0, 0, w, h, **kw)
|
|
634
|
+
|
|
635
|
+
def _render_content(self, ctx: RenderContext):
|
|
636
|
+
return None
|
|
637
|
+
|
|
638
|
+
|
|
639
|
+
def _as_grid(values, rows: int = None, cols: int = None):
|
|
640
|
+
if values is None:
|
|
641
|
+
return None
|
|
642
|
+
seq = list(values)
|
|
643
|
+
if not seq:
|
|
644
|
+
return [[]]
|
|
645
|
+
if isinstance(seq[0], (list, tuple)):
|
|
646
|
+
return [list(r) for r in seq]
|
|
647
|
+
if rows and cols:
|
|
648
|
+
return [list(seq[i * cols:(i + 1) * cols]) for i in range(rows)]
|
|
649
|
+
if cols:
|
|
650
|
+
return [list(seq[i:i + cols]) for i in range(0, len(seq), cols)]
|
|
651
|
+
if rows:
|
|
652
|
+
per = math.ceil(len(seq) / rows)
|
|
653
|
+
return [list(seq[i:i + per]) for i in range(0, len(seq), per)]
|
|
654
|
+
return [list(seq)]
|
|
655
|
+
|
|
656
|
+
|
|
657
|
+
class LabelledMatrix(Component):
|
|
658
|
+
"""A :class:`Matrix` with its axis labels, delimiters and caption.
|
|
659
|
+
|
|
660
|
+
The pieces of a matrix in a formula — a rotated row label, a column label,
|
|
661
|
+
brackets, a caption underneath — are each trivial on their own and fiddly
|
|
662
|
+
to place together. This bundles them into one movable unit.
|
|
663
|
+
|
|
664
|
+
>>> LabelledMatrix(values, row_label="seq len", col_label="$d$",
|
|
665
|
+
... caption="$\\Pi_{\\mathcal{NM}}$", brackets="round")
|
|
666
|
+
|
|
667
|
+
Exposes ``.matrix``, ``.caption_text``, ``.row_text`` and ``.col_text``.
|
|
668
|
+
"""
|
|
669
|
+
|
|
670
|
+
role = "group"
|
|
671
|
+
|
|
672
|
+
def build(self, values=None, *, cell=16, row_label=None, col_label=None,
|
|
673
|
+
caption=None, brackets=None, label_gap: float = 7.0,
|
|
674
|
+
caption_gap: float = 9.0, bracket_gap: float = 5.0,
|
|
675
|
+
label_size: float = 11.0, caption_size: float = 13.0,
|
|
676
|
+
label_style=None, **matrix_kw):
|
|
677
|
+
matrix = Matrix(values, cell=cell, add=False, **matrix_kw)
|
|
678
|
+
self.expose("matrix", matrix)
|
|
679
|
+
parts = [matrix]
|
|
680
|
+
|
|
681
|
+
if brackets:
|
|
682
|
+
parts.extend(self._delimiters(matrix, brackets, bracket_gap))
|
|
683
|
+
span = BBox.union_all([p.bbox for p in parts]) or matrix.bbox
|
|
684
|
+
|
|
685
|
+
if col_label is not None:
|
|
686
|
+
top = Text(col_label, add=False, font_size=label_size,
|
|
687
|
+
style=label_style)
|
|
688
|
+
top.at(span.cx, span.y0 - label_gap, anchor="s")
|
|
689
|
+
self.expose("col_text", top)
|
|
690
|
+
parts.append(top)
|
|
691
|
+
if row_label is not None:
|
|
692
|
+
side = Text(row_label, add=False, font_size=label_size,
|
|
693
|
+
style=label_style)
|
|
694
|
+
side.rotate(-90)
|
|
695
|
+
side.at(span.x0 - label_gap, span.cy, anchor="e")
|
|
696
|
+
self.expose("row_text", side)
|
|
697
|
+
parts.append(side)
|
|
698
|
+
if caption is not None:
|
|
699
|
+
below = Text(caption, add=False, font_size=caption_size)
|
|
700
|
+
below.at(span.cx, span.y1 + caption_gap, anchor="n")
|
|
701
|
+
self.expose("caption_text", below)
|
|
702
|
+
parts.append(below)
|
|
703
|
+
return parts
|
|
704
|
+
|
|
705
|
+
def _delimiters(self, matrix, kind, gap: float) -> list:
|
|
706
|
+
"""Square or round brackets hugging the matrix."""
|
|
707
|
+
box = matrix.bbox.expand(gap)
|
|
708
|
+
arm = min(box.w * 0.16, 7.0)
|
|
709
|
+
style = str(kind).lower()
|
|
710
|
+
stroke = dict(fill="none", stroke=self.prop("color", "#16181d"),
|
|
711
|
+
stroke_width=1.2, stroke_linecap="round",
|
|
712
|
+
stroke_linejoin="round", add=False)
|
|
713
|
+
if style in ("round", "paren", "()"):
|
|
714
|
+
bulge = max(6.0, box.h * 0.14)
|
|
715
|
+
left = (f"M{fmt(box.x0)} {fmt(box.y0)}"
|
|
716
|
+
f"Q{fmt(box.x0 - bulge)} {fmt(box.cy)} "
|
|
717
|
+
f"{fmt(box.x0)} {fmt(box.y1)}")
|
|
718
|
+
right = (f"M{fmt(box.x1)} {fmt(box.y0)}"
|
|
719
|
+
f"Q{fmt(box.x1 + bulge)} {fmt(box.cy)} "
|
|
720
|
+
f"{fmt(box.x1)} {fmt(box.y1)}")
|
|
721
|
+
else:
|
|
722
|
+
left = (f"M{fmt(box.x0 + arm)} {fmt(box.y0)}H{fmt(box.x0)}"
|
|
723
|
+
f"V{fmt(box.y1)}H{fmt(box.x0 + arm)}")
|
|
724
|
+
right = (f"M{fmt(box.x1 - arm)} {fmt(box.y0)}H{fmt(box.x1)}"
|
|
725
|
+
f"V{fmt(box.y1)}H{fmt(box.x1 - arm)}")
|
|
726
|
+
from .shapes import Path
|
|
727
|
+
return [Path(left, **stroke), Path(right, **stroke)]
|