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/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)]