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/frame.py ADDED
@@ -0,0 +1,376 @@
1
+ """Coordinate frames: map data space to figure space.
2
+
3
+ A :class:`Frame` is a group with a rectangular plot area and a data-to-world
4
+ mapping. ``frame.pt(x, y)`` turns a data coordinate into a figure point, so
5
+ you can mix hand-placed graphics and data-driven ones freely.
6
+
7
+ The plotting helpers are deliberately thin: they build ordinary figkit
8
+ elements you can restyle, anchor to and move afterwards.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import math
14
+
15
+ from .colors import colormap
16
+ from .core import Element, Group
17
+ from .geom import BBox, Point
18
+ from .shapes import Box, Line, Marker, Polygon, Polyline
19
+ from .text import Text
20
+
21
+ __all__ = ["Frame", "nice_ticks"]
22
+
23
+
24
+ def nice_ticks(lo: float, hi: float, n: int = 5) -> list:
25
+ """Human-friendly tick positions covering ``[lo, hi]``."""
26
+ if hi == lo:
27
+ return [lo]
28
+ if hi < lo:
29
+ lo, hi = hi, lo
30
+ raw = (hi - lo) / max(1, n)
31
+ mag = 10 ** math.floor(math.log10(raw))
32
+ for mult in (1, 2, 2.5, 5, 10):
33
+ if raw <= mult * mag:
34
+ step = mult * mag
35
+ break
36
+ else:
37
+ step = 10 * mag
38
+ start = math.ceil(lo / step) * step
39
+ out = []
40
+ v = start
41
+ while v <= hi + step * 1e-9:
42
+ out.append(round(v, 12) + 0.0)
43
+ v += step
44
+ return out
45
+
46
+
47
+ def _fmt_tick(v: float) -> str:
48
+ if v == int(v) and abs(v) < 1e15:
49
+ return str(int(v))
50
+ s = f"{v:.6g}"
51
+ return s
52
+
53
+
54
+ class Frame(Group):
55
+ """A data-space coordinate system with a rectangular plot area.
56
+
57
+ >>> fr = Frame(w=320, h=180, xlim=(0, 10), ylim=(0, 1))
58
+ >>> fr.axes(xlabel="epoch", ylabel="accuracy")
59
+ >>> fr.line(xs, ys, stroke="@primary", lw=2)
60
+ >>> arrow(box.e, fr.pt(6, 0.8)) # point at a data coordinate
61
+ """
62
+
63
+ role = "group"
64
+
65
+ def __init__(self, x: float = 0.0, y: float = 0.0, w: float = 300.0,
66
+ h: float = 200.0, *, xlim=(0.0, 1.0), ylim=(0.0, 1.0),
67
+ flip_y: bool = True, xscale: str = "linear",
68
+ yscale: str = "linear", background=None, border=None,
69
+ clip_data: bool = False, **kw):
70
+ super().__init__(**kw)
71
+ self.area = BBox(float(x), float(y), float(w), float(h))
72
+ self.xlim = (float(xlim[0]), float(xlim[1]))
73
+ self.ylim = (float(ylim[0]), float(ylim[1]))
74
+ self.flip_y = bool(flip_y)
75
+ self.xscale = xscale
76
+ self.yscale = yscale
77
+ self._bg = None
78
+ self._frame_box = None
79
+ if background is not None:
80
+ self._bg = Box(None, self.area.x, self.area.y, self.area.w,
81
+ self.area.h, fill=background, stroke="none",
82
+ padding=0, z=-100, add=False)
83
+ self.add(self._bg)
84
+ if border:
85
+ self._frame_box = Box(
86
+ None, self.area.x, self.area.y, self.area.w, self.area.h,
87
+ fill="none", stroke=border if isinstance(border, str) else "#6b7280",
88
+ padding=0, radius=0, z=50, add=False)
89
+ self.add(self._frame_box)
90
+ if clip_data:
91
+ self.clip = True
92
+
93
+ # -- the mapping -----------------------------------------------------
94
+ def _sx(self, v: float) -> float:
95
+ lo, hi = self.xlim
96
+ if self.xscale == "log":
97
+ lo, hi, v = math.log10(max(lo, 1e-12)), math.log10(max(hi, 1e-12)), \
98
+ math.log10(max(float(v), 1e-12))
99
+ return 0.5 if hi == lo else (float(v) - lo) / (hi - lo)
100
+
101
+ def _sy(self, v: float) -> float:
102
+ lo, hi = self.ylim
103
+ if self.yscale == "log":
104
+ lo, hi, v = math.log10(max(lo, 1e-12)), math.log10(max(hi, 1e-12)), \
105
+ math.log10(max(float(v), 1e-12))
106
+ return 0.5 if hi == lo else (float(v) - lo) / (hi - lo)
107
+
108
+ def px(self, x: float) -> float:
109
+ """Data x -> world x."""
110
+ return self.area.x0 + self._sx(x) * self.area.w
111
+
112
+ def py(self, y: float) -> float:
113
+ """Data y -> world y (flipped so larger values are higher)."""
114
+ t = self._sy(y)
115
+ return self.area.y1 - t * self.area.h if self.flip_y \
116
+ else self.area.y0 + t * self.area.h
117
+
118
+ def pt(self, x: float, y: float) -> Point:
119
+ """Data ``(x, y)`` -> world :class:`~figkit.geom.Point`."""
120
+ return Point(self.px(x), self.py(y))
121
+
122
+ def pts(self, xs, ys=None) -> list:
123
+ if ys is None:
124
+ return [self.pt(p[0], p[1]) for p in xs]
125
+ return [self.pt(x, y) for x, y in zip(xs, ys)]
126
+
127
+ def data(self, px: float, py: float) -> tuple:
128
+ """World point -> data coordinates (inverse mapping)."""
129
+ tx = (px - self.area.x0) / self.area.w if self.area.w else 0.0
130
+ ty = ((self.area.y1 - py) if self.flip_y else (py - self.area.y0))
131
+ ty = ty / self.area.h if self.area.h else 0.0
132
+ lo, hi = self.xlim
133
+ x = lo + tx * (hi - lo) if self.xscale != "log" else \
134
+ 10 ** (math.log10(lo) + tx * (math.log10(hi) - math.log10(lo)))
135
+ lo, hi = self.ylim
136
+ y = lo + ty * (hi - lo) if self.yscale != "log" else \
137
+ 10 ** (math.log10(lo) + ty * (math.log10(hi) - math.log10(lo)))
138
+ return x, y
139
+
140
+ @property
141
+ def plot_area(self) -> BBox:
142
+ return self.area
143
+
144
+ def clip_bbox(self) -> BBox:
145
+ """``clip_data=True`` clips to the plot area, not the axes and labels."""
146
+ return self.area
147
+
148
+ def move(self, dx: float = 0.0, dy: float = 0.0) -> "Frame":
149
+ """Move the frame: children *and* the data-space plot area."""
150
+ if dx == 0 and dy == 0:
151
+ return self
152
+ self.area = self.area.translated(dx, dy)
153
+ return super().move(dx, dy)
154
+
155
+ @property
156
+ def bbox(self) -> BBox:
157
+ bb = super().bbox
158
+ return bb.union(self.area) if len(self._children) else self.area
159
+
160
+ def autoscale(self, xs=None, ys=None, pad: float = 0.05) -> "Frame":
161
+ """Set limits from data with a little breathing room."""
162
+ if xs is not None and len(xs):
163
+ lo, hi = min(xs), max(xs)
164
+ m = (hi - lo) * pad if hi > lo else 1.0
165
+ self.xlim = (lo - m, hi + m)
166
+ if ys is not None and len(ys):
167
+ lo, hi = min(ys), max(ys)
168
+ m = (hi - lo) * pad if hi > lo else 1.0
169
+ self.ylim = (lo - m, hi + m)
170
+ return self
171
+
172
+ # -- helpers ---------------------------------------------------------
173
+ def _adopt(self, el):
174
+ self.add(el)
175
+ return el
176
+
177
+ def at_data(self, element: Element, x: float, y: float,
178
+ anchor: str = "center") -> Element:
179
+ """Position an existing element by data coordinates."""
180
+ p = self.pt(x, y)
181
+ element.at(p.x, p.y, anchor=anchor)
182
+ return self._adopt(element)
183
+
184
+ def text(self, label: str, x: float, y: float, anchor: str = "center",
185
+ **kw) -> Text:
186
+ t = Text(label, add=False, **kw)
187
+ p = self.pt(x, y)
188
+ t.at(p.x, p.y, anchor=anchor)
189
+ return self._adopt(t)
190
+
191
+ # -- data marks ------------------------------------------------------
192
+ def line(self, xs, ys=None, *, smooth: bool = False, **kw) -> Polyline:
193
+ """A polyline through data points."""
194
+ pts = self.pts(xs, ys)
195
+ kw.setdefault("fill", "none")
196
+ kw.setdefault("stroke", "@primary")
197
+ kw.setdefault("stroke_width", 2.0)
198
+ kw.setdefault("stroke_linejoin", "round")
199
+ kw.setdefault("stroke_linecap", "round")
200
+ if smooth:
201
+ from .connectors import _catmull_rom
202
+ from .shapes import Path
203
+ return self._adopt(Path(_catmull_rom(pts), add=False, **kw))
204
+ return self._adopt(Polyline(pts, add=False, **kw))
205
+
206
+ def area_fill(self, xs, ys=None, base: float = None, **kw) -> Polygon:
207
+ """Filled area between a series and a baseline."""
208
+ pts = self.pts(xs, ys)
209
+ if not pts:
210
+ return None
211
+ base = self.ylim[0] if base is None else base
212
+ yb = self.py(base)
213
+ poly = [Point(pts[0].x, yb)] + pts + [Point(pts[-1].x, yb)]
214
+ kw.setdefault("fill", "@primary_soft")
215
+ kw.setdefault("stroke", "none")
216
+ return self._adopt(Polygon(poly, add=False, **kw))
217
+
218
+ def scatter(self, xs, ys=None, *, size=7.0, shape: str = "circle",
219
+ colors=None, values=None, cmap="viridis", vmin: float = None,
220
+ vmax: float = None, **kw) -> Group:
221
+ """Scatter markers. ``values`` + ``cmap`` colour-codes them."""
222
+ pts = self.pts(xs, ys)
223
+ kw.setdefault("fill", "@primary")
224
+ kw.setdefault("stroke", "none")
225
+ lo = min(values) if (values and vmin is None) else vmin
226
+ hi = max(values) if (values and vmax is None) else vmax
227
+ out = Group(add=False)
228
+ for i, p in enumerate(pts):
229
+ style = dict(kw)
230
+ if colors:
231
+ style["fill"] = colors[i % len(colors)]
232
+ elif values is not None and i < len(values):
233
+ t = 0.5 if hi == lo else (float(values[i]) - lo) / (hi - lo)
234
+ style["fill"] = colormap(cmap, t)
235
+ s = size[i] if isinstance(size, (list, tuple)) else size
236
+ out.add(Marker(p, s, shape, add=False, **style))
237
+ return self._adopt(out)
238
+
239
+ def bars(self, xs, heights, *, width: float = 0.7, base: float = 0.0,
240
+ colors=None, horizontal: bool = False, **kw) -> Group:
241
+ """Bar chart. ``width`` is in data units (or a fraction of the step)."""
242
+ out = Group(add=False)
243
+ step = 1.0
244
+ if len(xs) > 1:
245
+ step = min(abs(float(xs[i + 1]) - float(xs[i]))
246
+ for i in range(len(xs) - 1)) or 1.0
247
+ bw = width * step if width <= 1.0 else width
248
+ kw.setdefault("fill", "@primary")
249
+ kw.setdefault("stroke", "none")
250
+ for i, (x, hgt) in enumerate(zip(xs, heights)):
251
+ style = dict(kw)
252
+ if colors:
253
+ style["fill"] = colors[i % len(colors)]
254
+ if horizontal:
255
+ x0, x1 = self.px(base), self.px(hgt)
256
+ y0, y1 = self.py(float(x) - bw / 2), self.py(float(x) + bw / 2)
257
+ else:
258
+ x0, x1 = self.px(float(x) - bw / 2), self.px(float(x) + bw / 2)
259
+ y0, y1 = self.py(base), self.py(hgt)
260
+ bar = Box(None, min(x0, x1), min(y0, y1), abs(x1 - x0),
261
+ abs(y1 - y0), padding=0, add=False, **style)
262
+ out.add(bar)
263
+ return self._adopt(out)
264
+
265
+ def region(self, x0: float, x1: float, y0: float, y1: float, **kw) -> Box:
266
+ """A shaded rectangle in data coordinates."""
267
+ a, b = self.pt(x0, y0), self.pt(x1, y1)
268
+ kw.setdefault("fill", "#0000000f")
269
+ kw.setdefault("stroke", "none")
270
+ kw.setdefault("padding", 0)
271
+ return self._adopt(Box(None, min(a.x, b.x), min(a.y, b.y),
272
+ abs(b.x - a.x), abs(b.y - a.y), add=False, **kw))
273
+
274
+ def hline(self, y: float, **kw) -> Line:
275
+ kw.setdefault("stroke", "#9aa1ac")
276
+ return self._adopt(Line(self.pt(self.xlim[0], y),
277
+ self.pt(self.xlim[1], y), add=False, **kw))
278
+
279
+ def vline(self, x: float, **kw) -> Line:
280
+ kw.setdefault("stroke", "#9aa1ac")
281
+ return self._adopt(Line(self.pt(x, self.ylim[0]),
282
+ self.pt(x, self.ylim[1]), add=False, **kw))
283
+
284
+ # -- axes ------------------------------------------------------------
285
+ def gridlines(self, xticks=None, yticks=None, n: int = 5, **kw) -> Group:
286
+ """Light grid lines behind the data."""
287
+ kw.setdefault("stroke", "#e5e7eb")
288
+ kw.setdefault("stroke_width", 1.0)
289
+ g = Group(add=False, z=-50)
290
+ for v in (nice_ticks(*self.xlim, n) if xticks is None else xticks):
291
+ g.add(Line(self.pt(v, self.ylim[0]), self.pt(v, self.ylim[1]),
292
+ add=False, **kw))
293
+ for v in (nice_ticks(*self.ylim, n) if yticks is None else yticks):
294
+ g.add(Line(self.pt(self.xlim[0], v), self.pt(self.xlim[1], v),
295
+ add=False, **kw))
296
+ return self._adopt(g)
297
+
298
+ def xaxis(self, ticks=None, n: int = 5, labels=None, fmt=_fmt_tick,
299
+ tick_size: float = 5.0, label_gap: float = 5.0,
300
+ title: str = None, title_gap: float = 6.0, side: str = "bottom",
301
+ show_line: bool = True, font_size: float = None, **kw) -> Group:
302
+ """Draw the x axis with ticks and labels."""
303
+ kw.setdefault("stroke", "#6b7280")
304
+ kw.setdefault("stroke_width", 1.0)
305
+ g = Group(add=False)
306
+ y = self.area.y1 if side == "bottom" else self.area.y0
307
+ sign = 1 if side == "bottom" else -1
308
+ if show_line:
309
+ g.add(Line((self.area.x0, y), (self.area.x1, y), add=False, **kw))
310
+ vals = nice_ticks(*self.xlim, n) if ticks is None else list(ticks)
311
+ texts = labels if labels is not None else [fmt(v) for v in vals]
312
+ for v, lab in zip(vals, texts):
313
+ px = self.px(v)
314
+ if tick_size:
315
+ g.add(Line((px, y), (px, y + sign * tick_size), add=False, **kw))
316
+ if lab is not None and lab != "":
317
+ t = Text(str(lab), add=False, font_size=font_size or 11,
318
+ color="#4b5563")
319
+ t.at(px, y + sign * (tick_size + label_gap),
320
+ anchor="n" if side == "bottom" else "s")
321
+ g.add(t)
322
+ if title:
323
+ t = Text(title, add=False, font_size=font_size or 12)
324
+ gb = g.bbox
325
+ t.at(self.area.cx, (gb.y1 + title_gap) if side == "bottom"
326
+ else (gb.y0 - title_gap),
327
+ anchor="n" if side == "bottom" else "s")
328
+ g.add(t)
329
+ return self._adopt(g)
330
+
331
+ def yaxis(self, ticks=None, n: int = 5, labels=None, fmt=_fmt_tick,
332
+ tick_size: float = 5.0, label_gap: float = 5.0,
333
+ title: str = None, title_gap: float = 6.0, side: str = "left",
334
+ show_line: bool = True, font_size: float = None,
335
+ title_rotate: bool = True, **kw) -> Group:
336
+ """Draw the y axis with ticks and labels."""
337
+ kw.setdefault("stroke", "#6b7280")
338
+ kw.setdefault("stroke_width", 1.0)
339
+ g = Group(add=False)
340
+ x = self.area.x0 if side == "left" else self.area.x1
341
+ sign = -1 if side == "left" else 1
342
+ if show_line:
343
+ g.add(Line((x, self.area.y0), (x, self.area.y1), add=False, **kw))
344
+ vals = nice_ticks(*self.ylim, n) if ticks is None else list(ticks)
345
+ texts = labels if labels is not None else [fmt(v) for v in vals]
346
+ for v, lab in zip(vals, texts):
347
+ py = self.py(v)
348
+ if tick_size:
349
+ g.add(Line((x, py), (x + sign * tick_size, py), add=False, **kw))
350
+ if lab is not None and lab != "":
351
+ t = Text(str(lab), add=False, font_size=font_size or 11,
352
+ color="#4b5563", align="right" if side == "left" else "left")
353
+ t.at(x + sign * (tick_size + label_gap), py,
354
+ anchor="e" if side == "left" else "w")
355
+ g.add(t)
356
+ if title:
357
+ t = Text(title, add=False, font_size=font_size or 12)
358
+ gb = g.bbox
359
+ tx = (gb.x0 - title_gap) if side == "left" else (gb.x1 + title_gap)
360
+ t.at(tx, self.area.cy, anchor="e" if side == "left" else "w")
361
+ if title_rotate:
362
+ t.rotate(-90 if side == "left" else 90)
363
+ t.center_at(tx + (-t.bbox.w / 2 if side == "left"
364
+ else t.bbox.w / 2), self.area.cy)
365
+ g.add(t)
366
+ return self._adopt(g)
367
+
368
+ def axes(self, xlabel: str = None, ylabel: str = None, n: int = 5,
369
+ grid: bool = False, **kw) -> Group:
370
+ """Convenience: both axes (and optionally grid lines) in one call."""
371
+ out = Group(add=False)
372
+ if grid:
373
+ out.add(self.gridlines(n=n))
374
+ out.add(self.xaxis(n=n, title=xlabel, **kw))
375
+ out.add(self.yaxis(n=n, title=ylabel, **kw))
376
+ return self._adopt(out)