plot3 0.4.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.
plot3/build.py ADDED
@@ -0,0 +1,3948 @@
1
+ """Grammar → wire format: stats, layer specs, HTML document."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ import numpy as np
8
+ import pandas as pd
9
+
10
+ import copy
11
+ import os
12
+
13
+ from plot3.encode import encode_codes, encode_norm, pack_u16, pack_u32, pack_u8
14
+ from plot3.mathtext import split_math
15
+ from plot3.geoms import (
16
+ _Geom,
17
+ aes,
18
+ coord_3d,
19
+ coord_cartesian,
20
+ coord_equal,
21
+ coord_polar,
22
+ geom_col,
23
+ scale_colour_continuous,
24
+ )
25
+ from plot3.scales import ordered_levels, Scale, col_values, fmt_num, fmt_ticks, nice_ticks, resolution
26
+
27
+ from plot3.stats3d import isosurface_levels, regular_grid_mesh
28
+ from plot3.table import (
29
+ ColumnNotFound,
30
+ category_labels,
31
+ count_by,
32
+ filter_equal,
33
+ get_columns,
34
+ group_pieces,
35
+ has_column,
36
+ materialize_columns,
37
+ n_rows,
38
+ numeric_array,
39
+ require_columns,
40
+ subsample_rows,
41
+ unique_levels,
42
+ )
43
+ from plot3.themes import _CONT_PALETTES, _THEMES as THEMES
44
+
45
+
46
+ def copy_geom_with_density_n(geom: _Geom, n: int) -> _Geom:
47
+ out = copy.copy(geom)
48
+ out._density_n = int(n)
49
+ return out
50
+
51
+
52
+ def _boxplot_stats(values: np.ndarray, coef: float = 1.5):
53
+ """Tukey five-number box + outliers (ggplot2 / geom_boxplot default)."""
54
+ values = np.asarray(values, dtype=np.float64)
55
+ values = values[np.isfinite(values)]
56
+ if values.size == 0:
57
+ return None
58
+ q1, med, q3 = np.percentile(values, [25, 50, 75])
59
+ if coef <= 0:
60
+ return (
61
+ float(values.min()),
62
+ float(q1),
63
+ float(med),
64
+ float(q3),
65
+ float(values.max()),
66
+ np.asarray([], dtype=np.float64),
67
+ )
68
+ iqr = q3 - q1
69
+ lo_fence = q1 - coef * iqr
70
+ hi_fence = q3 + coef * iqr
71
+ inside = values[(values >= lo_fence) & (values <= hi_fence)]
72
+ ymin = float(inside.min()) if inside.size else float(q1)
73
+ ymax = float(inside.max()) if inside.size else float(q3)
74
+ outliers = values[(values < lo_fence) | (values > hi_fence)]
75
+ return ymin, float(q1), float(med), float(q3), ymax, outliers
76
+
77
+
78
+ def _kde_1d(
79
+ values: np.ndarray,
80
+ *,
81
+ n: int = 512,
82
+ adjust: float = 1.0,
83
+ ) -> tuple[np.ndarray, np.ndarray]:
84
+ """Gaussian KDE on a regular grid (Scott bandwidth × adjust)."""
85
+ x = np.asarray(values, dtype=np.float64)
86
+ x = x[np.isfinite(x)]
87
+ n = max(8, int(n))
88
+ if x.size == 0:
89
+ return np.asarray([], dtype=np.float64), np.asarray([], dtype=np.float64)
90
+ if x.size == 1:
91
+ center = float(x[0])
92
+ grid = np.linspace(center - 1.0, center + 1.0, n)
93
+ dens = np.exp(-0.5 * ((grid - center) / 0.2) ** 2)
94
+ dens /= dens.sum() * (grid[1] - grid[0])
95
+ return grid, dens
96
+ std = float(np.std(x, ddof=1)) or 1e-6
97
+ bw = max(1e-9, float(adjust) * 1.06 * std * (x.size ** (-0.2)))
98
+ lo = float(x.min()) - 3.0 * bw
99
+ hi = float(x.max()) + 3.0 * bw
100
+ if hi <= lo:
101
+ hi = lo + 1.0
102
+ grid = np.linspace(lo, hi, n)
103
+ # (n_grid, n_obs)
104
+ u = (grid[:, None] - x[None, :]) / bw
105
+ dens = np.exp(-0.5 * u * u).sum(axis=1) / (
106
+ x.size * bw * math.sqrt(2.0 * math.pi)
107
+ )
108
+ return grid, dens
109
+
110
+
111
+ def _as_discrete_x(table, xcol: str):
112
+ """Cast count-key column to ordered categories so the x scale is discrete.
113
+
114
+ ggplot2 users typically map ``factor(x)`` for bar charts; without that,
115
+ integer/numeric codes sit on a continuous axis (ticks at 5, 7, …). Count
116
+ bars are categories of *values*, so we draw them discretely while keeping
117
+ the original labels. Relative ``width`` still defaults to 0.9 (user-settable).
118
+
119
+ Returns a small pandas frame (count tables are tiny; categorical levels
120
+ preserve first-appearance order from ``count_by``).
121
+ """
122
+ from plot3.table import as_pandas
123
+
124
+ out = as_pandas(table).copy()
125
+ raw = out[xcol]
126
+ if isinstance(raw.dtype, pd.CategoricalDtype):
127
+ present = set(raw.dropna().tolist())
128
+ ordered = [c for c in raw.cat.categories if c in present]
129
+ ordered += [None] if raw.isna().any() else []
130
+ else:
131
+ ordered = ordered_levels(raw.tolist())
132
+ levels = ["NA" if v is None else str(v) for v in ordered]
133
+ out[xcol] = raw.map(lambda v: "NA" if pd.isna(v) else str(v))
134
+ out[xcol] = pd.Categorical(out[xcol], categories=levels, ordered=True)
135
+ return out
136
+
137
+
138
+ _HIST_STATS = {
139
+ "count": "count", "after_stat(count)": "count", "..count..": "count", "stat(count)": "count",
140
+ "density": "density", "after_stat(density)": "density", "..density..": "density",
141
+ "stat(density)": "density",
142
+ }
143
+
144
+
145
+ def _hist_stat(mapping: dict) -> str:
146
+ """aes(y=after_stat(density)) for a histogram on the density scale."""
147
+ y = mapping.get("y")
148
+ if y is None:
149
+ return "count"
150
+ key = str(y).replace(" ", "")
151
+ if key not in _HIST_STATS:
152
+ raise ValueError(
153
+ f"geom_histogram(aes(y={y!r})): a histogram's y is computed. Use "
154
+ 'aes(y="after_stat(density)") for the density scale, or leave y out for counts'
155
+ )
156
+ return _HIST_STATS[key]
157
+
158
+
159
+ def _density_scale(counts: np.ndarray, edges: np.ndarray) -> np.ndarray:
160
+ """Counts as a density: each bar's area is its share of the rows."""
161
+ total = float(np.sum(counts))
162
+ widths = np.diff(np.asarray(edges, dtype=np.float64))
163
+ if total <= 0:
164
+ return np.zeros_like(counts, dtype=np.float64)
165
+ return np.asarray(counts, dtype=np.float64) / (total * widths)
166
+
167
+
168
+ def _grouped_histogram(geom, data, xcol, group_col, edges, centers, closed, stat="count"):
169
+ """One histogram per group on shared bins, stacked like ggplot2.
170
+
171
+ Returns None for a continuous colour, which is not a grouping.
172
+ """
173
+ from plot3 import stat2d
174
+
175
+ frame = materialize_columns(data, [xcol, group_col])
176
+ kind, _codes, _cats = col_values(frame[group_col])
177
+ if kind != "cat":
178
+ return None
179
+ rows = []
180
+ for level in ordered_levels(frame[group_col].tolist()):
181
+ values = frame.loc[frame[group_col] == level, xcol].to_numpy(np.float64)
182
+ values = values[np.isfinite(values)]
183
+ if closed == "left":
184
+ counts = _hist_counts_left_closed(values, edges)
185
+ else:
186
+ counts, _ = np.histogram(values, bins=edges)
187
+ if stat == "density":
188
+ counts = _density_scale(counts, edges) # each group integrates to 1
189
+ for center, count in zip(centers, counts):
190
+ rows.append({"__x": float(center), "__y": float(count), group_col: level})
191
+ table = pd.DataFrame(rows)
192
+ proxy = _Geom(color=geom.const_color, alpha=geom.alpha)
193
+ proxy.kind = "col"
194
+ proxy.width = 1.0 # bars touch, as in a histogram
195
+ proxy.position = getattr(geom, "position", "stack")
196
+ out = stat2d.positioned_bars(proxy, {"x": "__x", "y": "__y", "color": group_col}, table)
197
+ out._axis_labels = {"x": xcol, "y": stat}
198
+ return out
199
+
200
+
201
+ def _freqpoly(geom, mapping: dict, data):
202
+ """geom_freqpoly: the histogram's counts at bin centres, as a line per
203
+ colour group, padded with a zero bin each side (ggplot2's pad = TRUE)."""
204
+ from plot3 import stat2d
205
+
206
+ if "x" not in mapping:
207
+ raise ValueError("geom_freqpoly() requires aes(x=)")
208
+ xcol = mapping["x"]
209
+ stat = _hist_stat(mapping)
210
+ values = numeric_array(data, xcol, dropna=True)
211
+ if values.size == 0:
212
+ raise ValueError(f"geom_freqpoly(): no numeric values in {xcol!r}")
213
+ edges = _histogram_breaks(
214
+ values, bins=getattr(geom, "bins", None), binwidth=getattr(geom, "binwidth", None),
215
+ boundary=getattr(geom, "boundary", None), method=getattr(geom, "method", "fd"),
216
+ )
217
+ width = float(np.median(np.diff(edges)))
218
+ centres = 0.5 * (edges[:-1] + edges[1:])
219
+ xs = np.concatenate([[centres[0] - width], centres, [centres[-1] + width]])
220
+ colour = mapping.get("color")
221
+ colour = colour if colour and colour != xcol and has_column(data, colour) else None
222
+ frame = materialize_columns(data, [xcol] + ([colour] if colour else []))
223
+ if colour:
224
+ kind, _codes, _cats = col_values(frame[colour])
225
+ levels = ordered_levels(frame[colour].tolist()) if kind == "cat" else [None]
226
+ else:
227
+ levels = [None]
228
+ rows, starts = [], []
229
+ for level in levels:
230
+ part = frame if level is None else frame[frame[colour] == level]
231
+ vals = part[xcol].to_numpy(np.float64)
232
+ vals = vals[np.isfinite(vals)]
233
+ if getattr(geom, "closed", "right") == "left":
234
+ counts = _hist_counts_left_closed(vals, edges)
235
+ else:
236
+ counts, _ = np.histogram(vals, bins=edges)
237
+ if stat == "density":
238
+ counts = _density_scale(counts, edges)
239
+ ys = np.concatenate([[0.0], np.asarray(counts, dtype=np.float64), [0.0]])
240
+ starts.append([len(rows), len(xs)])
241
+ for x, y in zip(xs, ys):
242
+ rows.append({"x": float(x), "y": float(y), **({colour: level} if level is not None else {})})
243
+ out_frame = pd.DataFrame(rows)
244
+ out_map = {"x": "x", "y": "y", **({"colour": colour} if levels != [None] else {})}
245
+ out = stat2d._layer("line", out_frame, out_map, geom, _groups=starts,
246
+ linewidth=float(getattr(geom, "linewidth", None) or 2.0))
247
+ out._axis_labels = {"x": xcol, "y": stat}
248
+ return out
249
+
250
+
251
+ def _hist_counts_left_closed(values: np.ndarray, edges: np.ndarray) -> np.ndarray:
252
+ """Count with left-closed / right-open bins (ggplot2 ``closed = "left"``)."""
253
+ values = np.asarray(values, dtype=np.float64)
254
+ edges = np.asarray(edges, dtype=np.float64)
255
+ n = len(edges) - 1
256
+ # searchsorted side='left': edge[i] <= x < edge[i+1] is not default;
257
+ # use side='right' then subtract 1 for left-closed intervals.
258
+ idx = np.searchsorted(edges, values, side="right") - 1
259
+ idx = np.clip(idx, 0, n - 1)
260
+ # Drop values outside [edges[0], edges[-1]]
261
+ inside = (values >= edges[0]) & (values <= edges[-1])
262
+ counts = np.bincount(idx[inside], minlength=n).astype(np.float64)
263
+ return counts
264
+
265
+
266
+ def _edges_from_binwidth(
267
+ lo: float,
268
+ hi: float,
269
+ w: float,
270
+ boundary: float | None,
271
+ ) -> np.ndarray:
272
+ """Build equal-width edges covering [lo, hi] with optional boundary align."""
273
+ if w <= 0 or not math.isfinite(w):
274
+ raise ValueError("binwidth must be a positive finite number")
275
+ if boundary is not None:
276
+ origin = float(boundary)
277
+ start = origin + math.floor((lo - origin) / w) * w
278
+ else:
279
+ start = lo
280
+ n = max(1, int(math.ceil((hi - start) / w)))
281
+ edges = start + np.arange(n + 1, dtype=np.float64) * w
282
+ if edges[-1] < hi:
283
+ edges = np.append(edges, edges[-1] + w)
284
+ return edges
285
+
286
+
287
+ def _auto_binwidth(values: np.ndarray, method: str = "fd") -> float:
288
+ """Data-driven bin width (Freedman–Diaconis by default).
289
+
290
+ Falls back along Scott → Sturges-derived width when a rule is undefined
291
+ (e.g. zero IQR or zero variance).
292
+ """
293
+ x = np.asarray(values, dtype=np.float64)
294
+ x = x[np.isfinite(x)]
295
+ n = int(x.size)
296
+ if n <= 1:
297
+ return 1.0
298
+ lo = float(np.min(x))
299
+ hi = float(np.max(x))
300
+ span = hi - lo
301
+ if span <= 0:
302
+ return 1.0
303
+
304
+ method = (method or "fd").lower()
305
+ n_root = n ** (1.0 / 3.0)
306
+
307
+ def fd_width() -> float | None:
308
+ q75, q25 = np.percentile(x, [75.0, 25.0])
309
+ iqr = float(q75 - q25)
310
+ if iqr > 0:
311
+ return 2.0 * iqr / n_root
312
+ return None
313
+
314
+ def scott_width() -> float | None:
315
+ # Sample SD; Scott's normal-reference rule.
316
+ sd = float(np.std(x, ddof=1)) if n > 1 else 0.0
317
+ if sd > 0:
318
+ return 3.49 * sd / n_root
319
+ return None
320
+
321
+ def sturges_width() -> float:
322
+ # Convert Sturges bin count into a width over the data span.
323
+ k = max(1, int(math.ceil(math.log2(n) + 1.0)))
324
+ return span / k
325
+
326
+ order: list[str]
327
+ if method in {"fd", "freedman-diaconis", "freedman_diaconis"}:
328
+ order = ["fd", "scott", "sturges"]
329
+ elif method == "scott":
330
+ order = ["scott", "fd", "sturges"]
331
+ elif method == "sturges":
332
+ order = ["sturges"]
333
+ elif method == "auto":
334
+ # Prefer FD, then Scott, then Sturges (same cascade as a robust default).
335
+ order = ["fd", "scott", "sturges"]
336
+ else:
337
+ # Unknown name: still try FD cascade rather than failing hard here;
338
+ # numpy edges path handles named rules when bins is the method string.
339
+ order = ["fd", "scott", "sturges"]
340
+
341
+ for name in order:
342
+ if name == "fd":
343
+ w = fd_width()
344
+ elif name == "scott":
345
+ w = scott_width()
346
+ else:
347
+ w = sturges_width()
348
+ if w is not None and w > 0 and math.isfinite(w):
349
+ # At least one bin, at most a fine grid (avoid pathological widths).
350
+ return float(min(max(w, span / 1000.0), span))
351
+ return sturges_width()
352
+
353
+
354
+ def _histogram_breaks(
355
+ values: np.ndarray,
356
+ *,
357
+ bins: int | None,
358
+ binwidth: float | None,
359
+ boundary: float | None,
360
+ method: str = "fd",
361
+ ) -> np.ndarray:
362
+ """Compute histogram edges.
363
+
364
+ Priority (user settings always win when given):
365
+
366
+ 1. ``binwidth`` — absolute width (optional ``boundary`` alignment)
367
+ 2. ``bins`` — explicit bin count (optional ``boundary``)
368
+ 3. automatic width from ``method`` (default Freedman–Diaconis), then edges
369
+ """
370
+ values = np.asarray(values, dtype=np.float64)
371
+ values = values[np.isfinite(values)]
372
+ if values.size == 0:
373
+ return np.asarray([0.0, 1.0], dtype=np.float64)
374
+ lo = float(np.min(values))
375
+ hi = float(np.max(values))
376
+ if hi <= lo:
377
+ # Single unique value: one bin of unit width around the point.
378
+ mid = lo
379
+ w = 1.0 if binwidth is None else float(binwidth)
380
+ return np.asarray([mid - 0.5 * w, mid + 0.5 * w], dtype=np.float64)
381
+
382
+ if binwidth is not None:
383
+ return _edges_from_binwidth(lo, hi, float(binwidth), boundary)
384
+
385
+ if bins is not None:
386
+ n_bins = max(1, int(bins))
387
+ if boundary is None:
388
+ return np.linspace(lo, hi, n_bins + 1)
389
+ w = (hi - lo) / n_bins
390
+ return _edges_from_binwidth(lo, hi, w, boundary)
391
+
392
+ # Automatic: prefer numpy's named rule when it matches; otherwise FD cascade.
393
+ method = (method or "fd").lower()
394
+ numpy_names = {
395
+ "fd",
396
+ "scott",
397
+ "sturges",
398
+ "auto",
399
+ "doane",
400
+ "stone",
401
+ "rice",
402
+ "sqrt",
403
+ }
404
+ if method in numpy_names and boundary is None:
405
+ try:
406
+ edges = np.histogram_bin_edges(values, bins=method)
407
+ edges = np.asarray(edges, dtype=np.float64)
408
+ if edges.size >= 2 and np.all(np.isfinite(edges)):
409
+ return edges
410
+ except Exception:
411
+ pass
412
+ # FD/Scott/Sturges width cascade (also used when boundary is set).
413
+ w = _auto_binwidth(values, method=method)
414
+ return _edges_from_binwidth(lo, hi, w, boundary)
415
+
416
+
417
+ # Filled shapes take their colour from aes(fill=). Points, lines, error bars,
418
+ # and text use aes(colour=) only (ggplot2's default shapes ignore fill).
419
+ _FILL_KINDS = frozenset({
420
+ "col", "bar", "histogram", "boxplot", "box", "violin", "poly", "polygon", "area",
421
+ "density", "surface", "isosurface", "density_3d_stat", "ribbon",
422
+ "tile", "rect", "area_stat",
423
+ })
424
+
425
+
426
+ def _apply_fill(mapping: dict, kind: str) -> dict:
427
+ """Resolve ``fill`` into the one colour channel a layer draws with."""
428
+ out = dict(mapping)
429
+ fill = out.pop("fill", None)
430
+ if fill is None:
431
+ return out
432
+ if kind in _FILL_KINDS:
433
+ out["color"] = fill
434
+ elif kind == "smooth" and "color" not in out:
435
+ out["color"] = fill # a smoother's band follows its fill
436
+ else:
437
+ # Not drawn in the fill colour, but a discrete fill still groups the
438
+ # layer (ggplot2), so error bars dodge onto the bars they belong to.
439
+ out["__fillgroup"] = fill
440
+ return out
441
+
442
+
443
+ _REF_KINDS = frozenset({"hline", "vline", "abline"})
444
+ _MM_TO_PX = 96.0 / 25.4
445
+
446
+
447
+ def _text_annotations(geom, vals, color_scale, theme, palette=None) -> list[dict]:
448
+ """geom_text / geom_label rows as annotation entries in scale units."""
449
+ xs = np.asarray(vals["x"], dtype=np.float64) + float(getattr(geom, "nudge_x", 0.0))
450
+ ys = np.asarray(vals["y"], dtype=np.float64) + float(getattr(geom, "nudge_y", 0.0))
451
+ labels = vals.get("label") or []
452
+ colours: list[str | None] = [geom.const_color or theme["ink"]] * len(labels)
453
+ if vals.get("color") and vals["color"][0] == "cat" and color_scale and color_scale[0] == "cat":
454
+ _kind, codes, local = vals["color"]
455
+ palette = palette or theme["cat"]
456
+ for i, code in enumerate(codes):
457
+ if math.isfinite(code):
458
+ name = local[int(code)]
459
+ index = color_scale[1].index(name) if name in color_scale[1] else 0
460
+ colours[i] = palette[index % len(palette)]
461
+ face = getattr(geom, "fontface", "plain")
462
+ style = {
463
+ "style": "label" if getattr(geom, "_box", False) else "text",
464
+ "size": float(getattr(geom, "size", 3.88)) * _MM_TO_PX,
465
+ "hjust": float(getattr(geom, "hjust", 0.5)),
466
+ "vjust": float(getattr(geom, "vjust", 0.5)),
467
+ "weight": 700 if "bold" in face else 400,
468
+ "italic": "italic" in face,
469
+ "overlap": not bool(getattr(geom, "check_overlap", False)),
470
+ "alpha": 1.0 if geom.alpha is None else float(geom.alpha),
471
+ }
472
+ out = []
473
+ for x, y, text, colour in zip(xs, ys, labels, colours):
474
+ if text and math.isfinite(x) and math.isfinite(y):
475
+ out.append({"x": float(x), "y": float(y), "text": text, "color": colour, **style})
476
+ return out
477
+
478
+
479
+ def _ref_scale_value(scale, value) -> float | None:
480
+ """A reference value in the scale's own units (index, seconds, log10)."""
481
+ if scale.kind == "cat":
482
+ key = str(value)
483
+ return float(scale.cats.index(key)) if key in scale.cats else None
484
+ if scale.kind == "dt":
485
+ try:
486
+ return float(pd.Timestamp(value).timestamp())
487
+ except (TypeError, ValueError):
488
+ return None
489
+ try:
490
+ number = float(value)
491
+ except (TypeError, ValueError):
492
+ return None
493
+ if getattr(scale, "trans", None) == "log10":
494
+ return math.log10(number) if number > 0 else None
495
+ return number if math.isfinite(number) else None
496
+
497
+
498
+ def _ref_specs(ref_layers, scales, theme) -> list[dict]:
499
+ from plot3.geoms import dash_pattern
500
+
501
+ out = []
502
+ for ref in ref_layers:
503
+ style = {
504
+ "color": _hex_or_none(ref.const_color) or theme["ink"],
505
+ "width": float(getattr(ref, "linewidth", 1.0) or 1.0),
506
+ "alpha": 1.0 if ref.alpha is None else float(ref.alpha),
507
+ "dash": list(dash_pattern(getattr(ref, "linetype", None)) or []) or None,
508
+ }
509
+ if ref.kind in {"hline", "vline"}:
510
+ axis = "y" if ref.kind == "hline" else "x"
511
+ if axis not in scales:
512
+ continue
513
+ for value in ref.values:
514
+ v = _ref_scale_value(scales[axis], value)
515
+ if v is not None:
516
+ out.append({"kind": ref.kind, "value": v, **style})
517
+ else:
518
+ sx, sy = scales.get("x"), scales.get("y")
519
+ linear = all(
520
+ sc is not None and sc.kind == "num" and getattr(sc, "trans", None) is None
521
+ for sc in (sx, sy)
522
+ )
523
+ if not linear:
524
+ raise ValueError("geom_abline() needs numeric, linear x and y axes")
525
+ out.append({
526
+ "kind": "abline", "slope": ref.slope, "intercept": ref.intercept, **style,
527
+ })
528
+ return out
529
+
530
+
531
+ _STAT2D_KINDS = frozenset(
532
+ {"jitter", "errorbar", "linerange", "pointrange", "ribbon", "smooth", "summary",
533
+ "tile", "area_stat", "step", "segment", "rect", "qq", "qq_line", "ecdf",
534
+ "crossbar", "errorbarh", "polygon", "bin_2d", "hex", "density_2d",
535
+ "density_2d_filled", "contour", "ellipse", "count"}
536
+ )
537
+
538
+
539
+ def expand_stat_geom(
540
+ geom: _Geom,
541
+ base_mapping: aes,
542
+ data,
543
+ domains: dict | None = None,
544
+ transition=None,
545
+ slider=None,
546
+ coord=None,
547
+ addons=None,
548
+ ):
549
+ """Turn statistical geoms into concrete drawable layers.
550
+
551
+ Stats run against the native table backend (pandas / polars / tidy→polars).
552
+ Only selected columns are pulled to arrays; small *result* frames used as
553
+ ``data_override`` are plain pandas (already computed, render-ready).
554
+ Identity geoms (point / line / …) are returned unchanged and materialise
555
+ columns at the render boundary.
556
+
557
+ A formula with an integral, tangent, or inequality returns a list of
558
+ layers. A plain curve still returns one geom.
559
+ """
560
+ kind = getattr(geom, "kind", None)
561
+ if kind == "function":
562
+ from plot3.function import expand_function
563
+
564
+ layers = expand_function(
565
+ geom, base_mapping, data, domains, transition, slider, coord, addons
566
+ )
567
+ return layers[0] if len(layers) == 1 else layers
568
+ if kind == "vector":
569
+ from plot3.calculus import expand_vector_field
570
+
571
+ layers = expand_vector_field(geom, transition, slider)
572
+ return layers[0] if len(layers) == 1 else layers
573
+ # Resolve fill per mapping, so a layer's own colour beats a base fill.
574
+ mapping = _apply_fill(dict(base_mapping), geom.kind)
575
+ mapping.update(_apply_fill(dict(geom.mapping), geom.kind))
576
+ if geom.kind == "box3d":
577
+ from plot3.stats3d import box3d_layers
578
+
579
+ return box3d_layers(geom, mapping, data)
580
+ if geom.kind == "blank":
581
+ # geom_blank: an invisible point layer, so its data train the scales.
582
+ out = copy.copy(geom)
583
+ out.kind = "point"
584
+ out._blank = True
585
+ out.position = "identity"
586
+ return out
587
+ if geom.kind == "point" and getattr(geom, "position", None) not in (None, "identity"):
588
+ # position_jitter(), position_jitterdodge(), position_nudge().
589
+ if "z" in mapping:
590
+ raise ValueError("geom_point(position=) moves points in 2D figures only")
591
+ from plot3 import stat2d
592
+
593
+ return stat2d.jitter(geom, mapping, data)
594
+ if geom.kind in _STAT2D_KINDS or geom.kind in {"col", "bar"}:
595
+ from plot3 import stat2d
596
+
597
+ if geom.kind in {"col", "bar"}:
598
+ if stat2d.wants_positioned_bars(geom, mapping, data):
599
+ return stat2d.positioned_bars(geom, mapping, data)
600
+ elif geom.kind == "jitter":
601
+ return stat2d.jitter(geom, mapping, data)
602
+ elif geom.kind in {"errorbar", "linerange", "pointrange"}:
603
+ return stat2d.ranges(geom, mapping, data)
604
+ elif geom.kind == "ribbon":
605
+ return stat2d.ribbon(geom, mapping, data)
606
+ elif geom.kind == "smooth":
607
+ return stat2d.smooth(geom, mapping, data)
608
+ elif geom.kind == "summary":
609
+ return stat2d.summary(geom, mapping, data)
610
+ else:
611
+ handler = {
612
+ "tile": stat2d.tile, "area_stat": stat2d.area, "step": stat2d.step,
613
+ "segment": stat2d.segment, "rect": stat2d.rect, "qq": stat2d.qq,
614
+ "qq_line": stat2d.qq_line, "ecdf": stat2d.ecdf,
615
+ "crossbar": stat2d.crossbar, "errorbarh": stat2d.errorbarh,
616
+ "polygon": stat2d.polygon, "bin_2d": stat2d.bin_2d,
617
+ "hex": stat2d.hex_bins, "density_2d": stat2d.density_2d,
618
+ "density_2d_filled": stat2d.density_2d, "contour": stat2d.contour,
619
+ "ellipse": stat2d.ellipse, "count": stat2d.count_points,
620
+ }[geom.kind]
621
+ return handler(geom, mapping, data)
622
+ if geom.kind == "bar":
623
+ if "x" not in mapping:
624
+ raise ValueError("geom_bar() requires aes(x=)")
625
+ xcol = mapping["x"]
626
+ # Backend-native count: pandas groupby / polars group_by / tidy→polars.
627
+ counts = count_by(data, xcol)
628
+ # ggplot2 stat_count at unique x; draw on a discrete scale (factor(x)
629
+ # style) so integer/numeric codes don't sit on a continuous axis with
630
+ # phantom ticks. Labels keep the original values as strings.
631
+ counts = _as_discrete_x(counts, xcol)
632
+ out = geom_col(
633
+ aes(x=xcol, y="y"),
634
+ width=getattr(geom, "width", 0.9),
635
+ color=geom.const_color,
636
+ colour=None,
637
+ alpha=geom.alpha,
638
+ )
639
+ out.data_override = counts
640
+ out.const_color = geom.const_color
641
+ out.alpha = geom.alpha
642
+ out._axis_labels = {"y": "count"}
643
+ return out
644
+ if geom.kind == "freqpoly":
645
+ return _freqpoly(geom, mapping, data)
646
+ if geom.kind == "histogram":
647
+ if "x" not in mapping:
648
+ raise ValueError("geom_histogram() requires aes(x=)")
649
+ xcol = mapping["x"]
650
+ stat = _hist_stat(mapping)
651
+ values = numeric_array(data, xcol, dropna=True)
652
+ binwidth = getattr(geom, "binwidth", None)
653
+ bins = getattr(geom, "bins", None)
654
+ method = getattr(geom, "method", "fd")
655
+ boundary = getattr(geom, "boundary", None)
656
+ closed = getattr(geom, "closed", "right")
657
+ if values.size == 0:
658
+ frame = pd.DataFrame(
659
+ {"x": np.array([], dtype=float), "y": np.array([], dtype=float)}
660
+ )
661
+ abs_width = 1.0
662
+ edge_lo, edge_hi = 0.0, 1.0
663
+ else:
664
+ edges = _histogram_breaks(
665
+ values,
666
+ bins=bins,
667
+ binwidth=binwidth,
668
+ boundary=boundary,
669
+ method=method,
670
+ )
671
+ # numpy: closed right by default (right edge included except last);
672
+ # closed left uses left-edge convention via density weights unused here.
673
+ # numpy histogram is right-closed (except last bin); ggplot2
674
+ # closed="right" matches that convention.
675
+ if closed == "left":
676
+ # Shift values slightly so membership matches left-closed bins
677
+ # without changing edges: count with inverted edges via
678
+ # searchsorted-based assignment.
679
+ counts = _hist_counts_left_closed(values, edges)
680
+ else:
681
+ counts, _ = np.histogram(values, bins=edges)
682
+ centers = 0.5 * (edges[:-1] + edges[1:])
683
+ abs_width = (
684
+ float(np.median(np.diff(edges))) if len(edges) > 1 else 1.0
685
+ )
686
+ edge_lo, edge_hi = float(edges[0]), float(edges[-1])
687
+ group_col = mapping.get("color")
688
+ if group_col and group_col != xcol and has_column(data, group_col):
689
+ grouped = _grouped_histogram(
690
+ geom, data, xcol, group_col, edges, centers, closed, stat
691
+ )
692
+ if grouped is not None:
693
+ return grouped
694
+ if stat == "density":
695
+ counts = _density_scale(counts, edges)
696
+ frame = pd.DataFrame(
697
+ {"x": centers, "y": counts.astype(np.float64)}
698
+ )
699
+ # Absolute bin width → bars touch (ggplot2 geom_histogram / GeomBar).
700
+ # Relative width stays 1.0 (full bin); users change bins/binwidth, not
701
+ # a gap fraction, for histograms.
702
+ out = geom_col(
703
+ aes(x="x", y="y"),
704
+ width=1.0,
705
+ color=geom.const_color,
706
+ alpha=geom.alpha,
707
+ )
708
+ out.data_override = frame
709
+ out.const_color = geom.const_color
710
+ out.alpha = geom.alpha
711
+ out._bar_width_data = abs_width # absolute data units (= binwidth)
712
+ out._axis_labels = {"x": xcol, "y": stat}
713
+ # Expand continuous domain to full bin edges (not just centres).
714
+ out._x_domain = (edge_lo, edge_hi)
715
+ return out
716
+ if geom.kind == "boxplot":
717
+ if "x" not in mapping or "y" not in mapping:
718
+ raise ValueError("geom_boxplot() requires aes(x=, y=)")
719
+ xcol, ycol = mapping["x"], mapping["y"]
720
+ require_columns(data, [xcol, ycol])
721
+ colour_col = mapping.get("color")
722
+ group_cols = [xcol]
723
+ if colour_col and colour_col != xcol and has_column(data, colour_col):
724
+ group_cols.append(colour_col)
725
+ coef = float(getattr(geom, "coef", 1.5))
726
+ rows: list[dict] = []
727
+ outlier_rows: list[dict] = []
728
+ for key_tuple, piece in group_pieces(data, group_cols):
729
+ values = numeric_array(piece, ycol, dropna=False)
730
+ stats = _boxplot_stats(values, coef=coef)
731
+ if stats is None:
732
+ continue
733
+ ymin, lower, middle, upper, ymax, outliers = stats
734
+ row = {
735
+ xcol: key_tuple[0],
736
+ "ymin": ymin,
737
+ "lower": lower,
738
+ "middle": middle,
739
+ "upper": upper,
740
+ "ymax": ymax,
741
+ }
742
+ if len(group_cols) > 1:
743
+ row[colour_col] = key_tuple[1]
744
+ rows.append(row)
745
+ for value in outliers:
746
+ out_row = {xcol: key_tuple[0], ycol: float(value)}
747
+ if len(group_cols) > 1:
748
+ out_row[colour_col] = key_tuple[1]
749
+ outlier_rows.append(out_row)
750
+ frame = pd.DataFrame(rows)
751
+ box_width = float(getattr(geom, "width", 0.75))
752
+ dodge_levels = None
753
+ if len(group_cols) > 1 and not frame.empty:
754
+ # ggplot2 dodges boxes by a second grouping: F and M side by side
755
+ # within each arm, sharing the 0.75 slot.
756
+ from plot3.stat2d import _axis, _dodge_offsets
757
+
758
+ kind_c, _codes, _cats = col_values(frame[colour_col])
759
+ if kind_c == "cat":
760
+ x_axis = _axis(materialize_columns(data, [xcol])[xcol])
761
+ if x_axis.kind == "cat" and x_axis.levels is not None:
762
+ x_index = {str(v): i for i, v in enumerate(x_axis.levels)}
763
+ c_levels = [str(v) for v in ordered_levels(frame[colour_col].tolist())]
764
+ c_index = {v: i for i, v in enumerate(c_levels)}
765
+ if len(c_levels) > 1:
766
+ offsets = _dodge_offsets(len(c_levels), box_width)
767
+
768
+ def dodged(x_value, c_value):
769
+ return x_index[str(x_value)] + offsets[c_index[str(c_value)]]
770
+
771
+ frame[xcol] = [dodged(a, b) for a, b in zip(frame[xcol], frame[colour_col])]
772
+ for row in outlier_rows:
773
+ row[xcol] = dodged(row[xcol], row[colour_col])
774
+ # position_dodge2's padding: a gap between boxes.
775
+ padding = float(getattr(getattr(geom, "position", None), "padding", 0.1))
776
+ box_width = box_width / len(c_levels) * (1.0 - padding)
777
+ dodge_levels = list(x_axis.levels)
778
+ if frame.empty:
779
+ frame = pd.DataFrame(
780
+ columns=[
781
+ xcol,
782
+ "ymin",
783
+ "lower",
784
+ "middle",
785
+ "upper",
786
+ "ymax",
787
+ *([colour_col] if colour_col and colour_col != xcol else []),
788
+ ]
789
+ )
790
+ # Concrete drawable layer with fixed stat column names.
791
+ out = _Geom(
792
+ aes(x=xcol, y="middle", colour=colour_col if colour_col else None),
793
+ color=geom.const_color,
794
+ alpha=geom.alpha,
795
+ )
796
+ out.kind = "box"
797
+ out.width = box_width
798
+ if dodge_levels is not None:
799
+ out._violin_levels = dodge_levels
800
+ out.outlier_size = float(getattr(geom, "outlier_size", 3.0))
801
+ out.data_override = frame
802
+ out.const_color = geom.const_color
803
+ out.alpha = geom.alpha
804
+ out._stat_y_cols = ("ymin", "lower", "middle", "upper", "ymax")
805
+ # geom_boxplot(outliers=False) or outlier_shape=None (ggplot2's
806
+ # outlier.shape = NA): no outlier points, for boxes under jitter.
807
+ hide = not getattr(geom, "outliers", True)
808
+ out._outlier_frame = pd.DataFrame([] if hide else outlier_rows)
809
+ out._y_name = ycol
810
+ raw = {**dict(base_mapping), **dict(geom.mapping)}
811
+ out._fill_mapped = "fill" in raw and "color" not in raw
812
+ return out
813
+ if geom.kind == "density":
814
+ if "x" not in mapping:
815
+ raise ValueError("geom_density() requires aes(x=)")
816
+ xcol = mapping["x"]
817
+ require_columns(data, [xcol])
818
+ colour_col = mapping.get("color")
819
+ n_grid = int(getattr(geom, "n", 512))
820
+ adjust = float(getattr(geom, "adjust", 1.0))
821
+ fill = bool(getattr(geom, "fill", True))
822
+ pieces: list[pd.DataFrame] = []
823
+ if colour_col and has_column(data, colour_col):
824
+ for key_tuple, piece in group_pieces(data, [colour_col]):
825
+ grid, dens = _kde_1d(
826
+ numeric_array(piece, xcol, dropna=False),
827
+ n=n_grid,
828
+ adjust=adjust,
829
+ )
830
+ if grid.size == 0:
831
+ continue
832
+ frame = pd.DataFrame(
833
+ {"x": grid, "y": dens, colour_col: key_tuple[0]}
834
+ )
835
+ pieces.append(frame)
836
+ else:
837
+ grid, dens = _kde_1d(
838
+ numeric_array(data, xcol, dropna=False),
839
+ n=n_grid,
840
+ adjust=adjust,
841
+ )
842
+ pieces.append(pd.DataFrame({"x": grid, "y": dens}))
843
+ frame = (
844
+ pd.concat(pieces, ignore_index=True)
845
+ if pieces
846
+ else pd.DataFrame(columns=["x", "y"])
847
+ )
848
+ map_kwargs = {"x": "x", "y": "y"}
849
+ if colour_col and colour_col in frame.columns:
850
+ map_kwargs["colour"] = colour_col
851
+ out = _Geom(
852
+ aes(**map_kwargs),
853
+ color=geom.const_color,
854
+ alpha=geom.alpha,
855
+ )
856
+ out.kind = "area" if fill else "line"
857
+ out.sort_x = True
858
+ out.linewidth = float(getattr(geom, "linewidth", 1.5))
859
+ out.data_override = frame
860
+ out.const_color = geom.const_color
861
+ out.alpha = geom.alpha if geom.alpha is not None else (0.35 if fill else 0.95)
862
+ out._baseline_zero = True
863
+ out._axis_labels = {"x": xcol, "y": "density"}
864
+ return out
865
+ if geom.kind == "violin":
866
+ if "x" not in mapping or "y" not in mapping:
867
+ raise ValueError("geom_violin() requires aes(x=, y=)")
868
+ xcol, ycol = mapping["x"], mapping["y"]
869
+ require_columns(data, [xcol, ycol])
870
+ colour_col = mapping.get("color")
871
+ n_grid = int(getattr(geom, "n", 128))
872
+ adjust = float(getattr(geom, "adjust", 1.0))
873
+ width = float(getattr(geom, "width", 0.9))
874
+ levels = category_labels(data, xcol)
875
+ level_index = {level: i for i, level in enumerate(levels)}
876
+ rows: list[dict] = []
877
+ grouping_cols = [xcol]
878
+ if colour_col and colour_col != xcol and has_column(data, colour_col):
879
+ grouping_cols.append(colour_col)
880
+ for key_tuple, piece in group_pieces(data, grouping_cols):
881
+ x_key = key_tuple[0]
882
+ x_pos = float(level_index.get(str(x_key), 0))
883
+ grid_y, dens = _kde_1d(
884
+ numeric_array(piece, ycol, dropna=False),
885
+ n=n_grid,
886
+ adjust=adjust,
887
+ )
888
+ if grid_y.size == 0:
889
+ continue
890
+ peak = float(np.nanmax(dens)) or 1.0
891
+ half = (dens / peak) * (width * 0.5)
892
+ # Closed polygon: left side bottom→top, right side top→bottom.
893
+ group_id = str(key_tuple)
894
+ for yv, hw in zip(grid_y, half):
895
+ row = {"x": x_pos - float(hw), "y": float(yv), "group": group_id}
896
+ if colour_col and colour_col != xcol:
897
+ row[colour_col] = key_tuple[1]
898
+ elif colour_col == xcol:
899
+ row[colour_col] = x_key
900
+ rows.append(row)
901
+ for yv, hw in zip(grid_y[::-1], half[::-1]):
902
+ row = {"x": x_pos + float(hw), "y": float(yv), "group": group_id}
903
+ if colour_col and colour_col != xcol:
904
+ row[colour_col] = key_tuple[1]
905
+ elif colour_col == xcol:
906
+ row[colour_col] = x_key
907
+ rows.append(row)
908
+ frame = pd.DataFrame(rows)
909
+ if frame.empty:
910
+ frame = pd.DataFrame(columns=["x", "y", "group"])
911
+ # Force categorical x scale labels via a helper frame for domains:
912
+ # encode x as numeric positions; stash category labels on the geom.
913
+ map_kwargs = {"x": "x", "y": "y", "group": "group"}
914
+ if colour_col and colour_col in frame.columns:
915
+ map_kwargs["colour"] = colour_col
916
+ out = _Geom(aes(**map_kwargs), color=geom.const_color, alpha=geom.alpha)
917
+ out.kind = "poly"
918
+ out.linewidth = float(getattr(geom, "linewidth", 1.0))
919
+ out.data_override = frame
920
+ out.const_color = geom.const_color
921
+ out.alpha = geom.alpha if geom.alpha is not None else 0.45
922
+ out._violin_levels = levels
923
+ out._is_violin = True
924
+ return out
925
+ if geom.kind == "surface":
926
+ if "x" not in mapping or "y" not in mapping or "z" not in mapping:
927
+ raise ValueError("geom_surface() requires aes(x=, y=, z=)")
928
+ xcol, ycol, zcol = mapping["x"], mapping["y"], mapping["z"]
929
+ ccol = mapping.get("color")
930
+ require_columns(data, [xcol, ycol, zcol])
931
+ # Only mesh columns cross the table→pandas boundary.
932
+ mesh_df = _surface_input_frame(data, xcol, ycol, zcol, ccol)
933
+ vertices, indices, nx, ny = regular_grid_mesh(
934
+ mesh_df, xcol, ycol, zcol, ccol=ccol
935
+ )
936
+ map_kwargs: dict = {"x": "x", "y": "y", "z": "z"}
937
+ if ccol and "colour" in vertices.columns:
938
+ map_kwargs["colour"] = "colour"
939
+ out = _Geom(
940
+ aes(**map_kwargs),
941
+ color=geom.const_color,
942
+ alpha=geom.alpha,
943
+ )
944
+ out.kind = "surface"
945
+ out.data_override = vertices
946
+ out.const_color = geom.const_color
947
+ out.alpha = geom.alpha if geom.alpha is not None else 0.95
948
+ out.wireframe = bool(getattr(geom, "wireframe", False))
949
+ out._indices = indices
950
+ out._nx = nx
951
+ out._ny = ny
952
+ return out
953
+ if geom.kind == "isosurface":
954
+ if "x" not in mapping or "y" not in mapping or "z" not in mapping:
955
+ raise ValueError("geom_isosurface() requires aes(x=, y=, z=)")
956
+ xcol, ycol, zcol = mapping["x"], mapping["y"], mapping["z"]
957
+ require_columns(data, [xcol, ycol, zcol])
958
+ # Column arrays only — no full-frame pandas conversion.
959
+ xs = numeric_array(data, xcol, dropna=False)
960
+ ys = numeric_array(data, ycol, dropna=False)
961
+ zs = numeric_array(data, zcol, dropna=False)
962
+ pts = np.column_stack([xs, ys, zs])
963
+ finite = np.isfinite(pts).all(axis=1)
964
+ pts = pts[finite]
965
+ n_bins = int(getattr(geom, "n", 32))
966
+ # Optional stat_density_3d on the figure is applied in build_spec.
967
+ n_bins = int(getattr(geom, "_density_n", n_bins))
968
+ levels = getattr(geom, "levels", (0.25, 0.5, 0.75))
969
+ vertices, indices, used = isosurface_levels(
970
+ pts, levels, n=n_bins, absolute=False
971
+ )
972
+ if vertices.empty:
973
+ vertices = pd.DataFrame(
974
+ columns=["x", "y", "z", "level", "colour"]
975
+ )
976
+ indices = np.zeros((0, 3), dtype=np.int32)
977
+ colour_by = getattr(geom, "colour_by", "level")
978
+ map_kwargs: dict = {"x": "x", "y": "y", "z": "z"}
979
+ if colour_by == "level" and "colour" in vertices.columns:
980
+ map_kwargs["colour"] = "colour"
981
+ out = _Geom(
982
+ aes(**map_kwargs),
983
+ color=geom.const_color,
984
+ alpha=geom.alpha,
985
+ )
986
+ out.kind = "isosurface"
987
+ out.data_override = vertices
988
+ out.const_color = geom.const_color
989
+ out.alpha = geom.alpha if geom.alpha is not None else 0.55
990
+ out.wireframe = bool(getattr(geom, "wireframe", False))
991
+ out._indices = indices
992
+ out._iso_levels = used
993
+ return out
994
+ return geom
995
+
996
+
997
+ def _surface_input_frame(data, xcol: str, ycol: str, zcol: str, ccol: str | None):
998
+ """Build a small pandas frame for ``regular_grid_mesh`` (xyz [+ colour]).
999
+
1000
+ Drops rows with non-finite x/y/z only; colour may remain null.
1001
+ Only the selected mesh columns are converted — not the full source table.
1002
+ """
1003
+ from plot3.table import _eager_polars, _polars_to_pandas, detect_backend
1004
+
1005
+ cols = [xcol, ycol, zcol]
1006
+ if ccol and has_column(data, ccol) and ccol not in cols:
1007
+ cols.append(ccol)
1008
+ backend = detect_backend(data)
1009
+ if backend == "pandas":
1010
+ work = data.loc[:, cols].copy()
1011
+ for c in (xcol, ycol, zcol):
1012
+ work[c] = pd.to_numeric(work[c], errors="coerce")
1013
+ return work.dropna(subset=[xcol, ycol, zcol])
1014
+
1015
+ import polars as pl
1016
+
1017
+ frame = _eager_polars(data).select(cols)
1018
+ for c in (xcol, ycol, zcol):
1019
+ frame = frame.with_columns(pl.col(c).cast(pl.Float64, strict=False))
1020
+ frame = frame.drop_nulls(subset=[xcol, ycol, zcol])
1021
+ return _polars_to_pandas(frame)
1022
+
1023
+
1024
+ # Geoms that cannot enter a 3D figure (stat expansions use these kinds too).
1025
+ _2D_ONLY_KINDS = frozenset(
1026
+ {"col", "box", "area", "poly", "bar", "histogram", "boxplot", "density", "violin"}
1027
+ )
1028
+ _3D_POINT_KINDS = frozenset({"point", "line", "surface", "isosurface"})
1029
+
1030
+
1031
+ def _default_3d_point_size(n: int, *, size_mode: str = "scene") -> float:
1032
+ """Point size in unit-cube scene units, smaller as the cloud gets denser.
1033
+
1034
+ The default camera shows the cube about 150 px across per scene unit
1035
+ on a 600 px viewer, so 0.035 is a 5 px mark for a few hundred points,
1036
+ 0.008 a fine 1.5 px grain for a 50,000-point lidar sweep, and 0.004 a
1037
+ 1 px grain for millions.
1038
+ """
1039
+ n = max(1, int(n))
1040
+ if size_mode == "screen":
1041
+ # Constant pixel size (Three.js sizeAttenuation=false).
1042
+ if n <= 5_000:
1043
+ return 2.0
1044
+ if n <= 50_000:
1045
+ return 1.5
1046
+ return 1.25
1047
+ # scene mode: world units after [0,1]×aspect encoding
1048
+ s = 0.03 * (1_000.0 / n) ** (1.0 / 3.0)
1049
+ return float(round(min(0.035, max(0.004, s)), 5))
1050
+
1051
+
1052
+ def _axis_label(g, base_map: dict, resolved, axis: str, is3d: bool) -> str:
1053
+ """Axis title: labs, then the ggplot mapping, then a function's variable.
1054
+
1055
+ ``labs(x="")`` removes the title (ggplot2's ``labs(x = NULL)``).
1056
+ """
1057
+ if (getattr(g, "theme_options", None) or {}).get(f"axis_title_{axis}") is False:
1058
+ return "" # theme(axis_title_x=element_blank())
1059
+ if axis in g.labs and g.labs[axis] is not None:
1060
+ return str(g.labs[axis])
1061
+ pscale = getattr(g, f"{axis}scale", None)
1062
+ if pscale is not None and pscale.name is not None:
1063
+ return str(pscale.name)
1064
+ mapped = base_map.get(axis)
1065
+ if mapped:
1066
+ # aes(y=after_stat(density)) is titled "density", as in ggplot2.
1067
+ return _HIST_STATS.get(str(mapped).replace(" ", ""), mapped)
1068
+ # A layer's own aes(y="v") names the axis before a computed layer's
1069
+ # suggestion (a ribbon's ymin, a smoother's y).
1070
+ for geom, mapping in resolved:
1071
+ if getattr(geom, "_replace_mapping", False) or getattr(geom, "_axis_labels", None):
1072
+ continue
1073
+ own = (getattr(geom, "mapping", None) or {}).get(axis)
1074
+ if own and own == mapping.get(axis):
1075
+ return str(own)
1076
+ for geom, _mapping in resolved:
1077
+ labels = getattr(geom, "_axis_labels", None)
1078
+ if labels and labels.get(axis):
1079
+ return str(labels[axis])
1080
+ if axis == "z" and not is3d:
1081
+ return ""
1082
+ return axis
1083
+
1084
+
1085
+ # A tall cloud (a helix, a tree) drawn to scale is a thin column in a wide
1086
+ # figure. Without coord_3d(aspect="data"), z is at most this many times the
1087
+ # wider horizontal side.
1088
+ _TALL_Z = 2.0
1089
+
1090
+
1091
+ def _auto_extent(scales) -> list[float] | None:
1092
+ """Box sides for aspect="auto", or None when true proportions fit."""
1093
+ spans = []
1094
+ for axis in ("x", "y", "z"):
1095
+ sc = scales.get(axis)
1096
+ lo, hi = getattr(sc, "lo", 0.0), getattr(sc, "hi", 1.0)
1097
+ try:
1098
+ span = abs(float(hi) - float(lo))
1099
+ except (TypeError, ValueError):
1100
+ span = 1.0
1101
+ spans.append(span if math.isfinite(span) and span > 0 else 1.0)
1102
+ across = max(spans[0], spans[1])
1103
+ if spans[2] <= _TALL_Z * across:
1104
+ return None
1105
+ sides = [spans[0], spans[1], _TALL_Z * across]
1106
+ top = max(sides)
1107
+ return [s / top for s in sides]
1108
+
1109
+
1110
+ def _coord_spec(coord, is3d: bool, resolved, scales=None) -> dict | None:
1111
+ """Coordinate spec for the viewer.
1112
+
1113
+ 3D keeps ``coord_3d``. A formula surface with no coord uses equal aspect
1114
+ (a cube); data such as lidar stays proportional. 2D uses ``coord_equal``
1115
+ when asked, and also when every layer is an implicit equation, so a
1116
+ circle is round in a wide panel.
1117
+ """
1118
+ if is3d:
1119
+ if isinstance(coord, coord_equal):
1120
+ raise ValueError(
1121
+ "coord_equal() is for 2D figures; use coord_3d(aspect='equal')"
1122
+ )
1123
+ if isinstance(coord, coord_polar):
1124
+ raise ValueError(
1125
+ "coord_polar() is for 2D figures. "
1126
+ 'For example ggplot() + geom_function("r = 1 + cos(theta)") '
1127
+ "+ coord_polar()"
1128
+ )
1129
+ if coord is not None:
1130
+ spec = coord.to_spec()
1131
+ if spec["aspect"] == "auto":
1132
+ ext = _auto_extent(scales or {})
1133
+ spec["aspect"] = "data" if ext is None else "auto"
1134
+ if ext is not None:
1135
+ spec["ext"] = ext
1136
+ return spec
1137
+ # A formula's axes are different quantities (t vs x); true proportions
1138
+ # can squash the surface to a sliver. Data stays proportional (lidar).
1139
+ if any(getattr(geom, "_function_surface", False) for geom, _ in resolved):
1140
+ return {"aspect": "equal", "sizeMode": "scene", "maxPoints": None}
1141
+ ext = _auto_extent(scales or {})
1142
+ if ext is not None:
1143
+ return {"aspect": "auto", "ext": ext, "sizeMode": "scene", "maxPoints": None}
1144
+ return {"aspect": "data", "sizeMode": "scene", "maxPoints": None}
1145
+ if isinstance(coord, coord_3d):
1146
+ raise ValueError(
1147
+ "coord_3d() requires a 3D figure (map aes(z=...) on layers)"
1148
+ )
1149
+ if isinstance(coord, coord_polar):
1150
+ return coord.to_spec()
1151
+ if isinstance(coord, coord_equal):
1152
+ return coord.to_spec()
1153
+ if isinstance(coord, coord_cartesian):
1154
+ return None if coord.expand else {"aspect": "data", "expand": False}
1155
+ if coord is not None:
1156
+ raise ValueError(
1157
+ "coord_3d() requires a 3D figure (map aes(z=...) on layers)"
1158
+ )
1159
+ implicit_only = bool(resolved) and all(
1160
+ getattr(geom, "_implicit", False) for geom, _mapped in resolved
1161
+ )
1162
+ if implicit_only:
1163
+ return {"aspect": "equal", "ratio": 1.0}
1164
+ return None
1165
+
1166
+
1167
+ def _log10_values(values: np.ndarray) -> tuple[np.ndarray, int]:
1168
+ """Map positive values to log10. Non-positive finites count as omitted."""
1169
+ v = np.asarray(values, dtype=np.float64)
1170
+ out = np.full(v.shape, np.nan, dtype=np.float64)
1171
+ ok = np.isfinite(v) & (v > 0)
1172
+ out[ok] = np.log10(v[ok])
1173
+ n_bad = int(np.count_nonzero(np.isfinite(v) & (v <= 0)))
1174
+ return out, n_bad
1175
+
1176
+
1177
+ def _last_finite(mat: np.ndarray) -> np.ndarray:
1178
+ """Last finite value of each row, or NaN when the row is all missing."""
1179
+ out = np.full(mat.shape[0], np.nan, dtype=np.float64)
1180
+ for frame in range(mat.shape[1] - 1, -1, -1):
1181
+ col = mat[:, frame]
1182
+ take = np.isnan(out) & np.isfinite(col)
1183
+ out[take] = col[take]
1184
+ return out
1185
+
1186
+
1187
+ def _size_breaks(vmax: float, integer: bool = False) -> list[float]:
1188
+ """Two to four round legend sizes up to ``vmax``, as ggplot2's breaks;
1189
+ whole numbers only when the data are counts."""
1190
+ if not math.isfinite(vmax) or vmax <= 0:
1191
+ return []
1192
+ exp = math.floor(math.log10(vmax)) - 1
1193
+ for shift in range(4):
1194
+ for mult in (1.0, 2.0, 2.5, 5.0):
1195
+ step = mult * 10.0 ** (exp + shift)
1196
+ if integer and abs(step - round(step)) > 1e-9:
1197
+ continue
1198
+ count = int(math.floor(vmax / step + 1e-9))
1199
+ if 2 <= count <= 4:
1200
+ return [float(f"{step * k:.12g}") for k in range(1, count + 1)]
1201
+ return [float(vmax)]
1202
+
1203
+
1204
+ def _area_fraction(values: np.ndarray, vmax: float) -> np.ndarray:
1205
+ """sqrt(value / vmax), so bubble area tracks the column. Missing stays NaN."""
1206
+ v = np.asarray(values, dtype=np.float64)
1207
+ out = np.full(v.shape, np.nan, dtype=np.float64)
1208
+ ok = np.isfinite(v) & (v >= 0)
1209
+ if vmax <= 0:
1210
+ out[ok] = 0.0
1211
+ else:
1212
+ out[ok] = np.sqrt(v[ok] / vmax)
1213
+ return out
1214
+
1215
+
1216
+ def _bubble_max(is3d: bool, coord) -> float:
1217
+ """Largest bubble diameter: pixels in 2D / screen mode, scene units in 3D."""
1218
+ if not is3d:
1219
+ return 46.0
1220
+ mode = "scene"
1221
+ if coord is not None:
1222
+ mode = getattr(coord, "size_mode", "scene") or "scene"
1223
+ if mode == "screen":
1224
+ return 32.0
1225
+ return 0.06
1226
+
1227
+
1228
+ def _filter_rows(vals: dict, ok: np.ndarray) -> None:
1229
+ """Keep rows where ``ok`` is true. Arrays and colour/group tuples follow."""
1230
+ n = int(ok.shape[0])
1231
+ for key, item in list(vals.items()):
1232
+ if isinstance(item, np.ndarray) and item.shape[:1] == (n,):
1233
+ vals[key] = item[ok]
1234
+ elif key == "color" and isinstance(item, tuple) and len(item) == 3:
1235
+ kind, cv, cats = item
1236
+ if isinstance(cv, np.ndarray) and cv.shape[:1] == (n,):
1237
+ vals[key] = (kind, cv[ok], cats)
1238
+ elif key == "group" and isinstance(item, tuple) and len(item) == 2:
1239
+ gv, gcats = item
1240
+ if isinstance(gv, np.ndarray) and gv.shape[:1] == (n,):
1241
+ vals[key] = (gv[ok], gcats)
1242
+ elif key == "ids" and isinstance(item, list) and len(item) == n:
1243
+ vals[key] = [item[i] for i, keep in enumerate(ok) if keep]
1244
+ elif key == "label" and isinstance(item, list) and len(item) == n:
1245
+ vals[key] = [item[i] for i, keep in enumerate(ok) if keep]
1246
+ elif key in {"shape", "linetype"} and isinstance(item, tuple) and len(item) == 2:
1247
+ codes, cats = item
1248
+ if isinstance(codes, np.ndarray) and codes.shape[:1] == (n,):
1249
+ vals[key] = (codes[ok], cats)
1250
+
1251
+
1252
+ def _pivot_transition(sub: pd.DataFrame, mapping: dict, transition) -> dict:
1253
+ """One row per group, and a matrix per channel shaped (objects, frames).
1254
+
1255
+ Matrices are still in data space. Object order is largest ``size`` first
1256
+ when that column is present, so small bubbles draw on top.
1257
+ """
1258
+ name = (
1259
+ "transition_states" if transition.kind == "states" else "transition_time"
1260
+ )
1261
+ if "group" not in mapping:
1262
+ raise ValueError(
1263
+ f"{name}() needs aes(group=) so each object keeps its identity "
1264
+ "across frames"
1265
+ )
1266
+ gcol = mapping["group"]
1267
+ tcol = transition.column
1268
+ if gcol not in sub.columns:
1269
+ raise KeyError(f"group column not in data: {gcol!r}")
1270
+ if tcol not in sub.columns:
1271
+ raise KeyError(f"{name}() column not in data: {tcol!r}")
1272
+
1273
+ g_ok = sub[gcol].notna().to_numpy()
1274
+ labels = sub[gcol].astype(str)
1275
+ ids = sorted({labels.iloc[i] for i in range(len(sub)) if g_ok[i]})
1276
+ if not ids:
1277
+ raise ValueError(f"{name}() needs at least one group value")
1278
+ id_index = {lab: i for i, lab in enumerate(ids)}
1279
+ g_idx = labels.map(id_index).to_numpy(dtype=np.float64)
1280
+ g_idx = np.where(g_ok, g_idx, -1).astype(np.int32)
1281
+
1282
+ if transition.kind == "states":
1283
+ t_ok = sub[tcol].notna().to_numpy()
1284
+ t_labels = sub[tcol].astype(str).tolist()
1285
+ times: list = []
1286
+ seen: dict[str, int] = {}
1287
+ for lab, ok in zip(t_labels, t_ok.tolist()):
1288
+ if not ok or lab in seen:
1289
+ continue
1290
+ seen[lab] = len(times)
1291
+ times.append(lab)
1292
+ if not times:
1293
+ raise ValueError(f"{name}() column has no values")
1294
+ t_idx = np.array(
1295
+ [seen.get(lab, -1) if ok else -1 for lab, ok in zip(t_labels, t_ok)],
1296
+ dtype=np.int32,
1297
+ )
1298
+ integer = False
1299
+ kind = "states"
1300
+ else:
1301
+ tvals = pd.to_numeric(sub[tcol], errors="coerce").to_numpy(dtype=np.float64)
1302
+ finite_t = np.isfinite(tvals)
1303
+ if not finite_t.any():
1304
+ raise ValueError(f"{name}() column has no finite values")
1305
+ times_arr = np.unique(tvals[finite_t])
1306
+ t_idx = np.full(len(sub), -1, dtype=np.int32)
1307
+ t_idx[finite_t] = np.searchsorted(times_arr, tvals[finite_t])
1308
+ scale = np.maximum(1.0, np.abs(times_arr))
1309
+ integer = bool(np.all(np.abs(times_arr - np.round(times_arr)) <= 1e-6 * scale))
1310
+ if integer:
1311
+ times = [int(round(float(t))) for t in times_arr]
1312
+ else:
1313
+ times = [float(t) for t in times_arr]
1314
+ kind = "time"
1315
+
1316
+ n_fr = len(times)
1317
+ valid = (g_idx >= 0) & (t_idx >= 0)
1318
+ if valid.any():
1319
+ key = g_idx[valid].astype(np.int64) * (n_fr + 1) + t_idx[valid]
1320
+ if np.unique(key).size != key.size:
1321
+ raise ValueError(
1322
+ f"{name}() found more than one row for the same group and frame"
1323
+ )
1324
+
1325
+ def _mat(values: np.ndarray) -> np.ndarray:
1326
+ out = np.full((len(ids), n_fr), np.nan, dtype=np.float64)
1327
+ if valid.any():
1328
+ out[g_idx[valid], t_idx[valid]] = np.asarray(values, dtype=np.float64)[valid]
1329
+ return out
1330
+
1331
+ order = np.arange(len(ids))
1332
+ return {
1333
+ "ids": ids,
1334
+ "times": times,
1335
+ "integer": integer,
1336
+ "kind": kind,
1337
+ "column": transition.column,
1338
+ "nFrames": n_fr,
1339
+ "g_idx": g_idx,
1340
+ "t_idx": t_idx,
1341
+ "valid": valid,
1342
+ "order": order,
1343
+ "_mat": _mat,
1344
+ }
1345
+
1346
+
1347
+ def _range_transition_meta(transition) -> dict:
1348
+ """Clock for a parameter sweep. Times follow the first range."""
1349
+ names = list(transition.ranges)
1350
+ first = names[0]
1351
+ lo, hi = transition.ranges[first]
1352
+ n_frames = int(transition.frames)
1353
+ times_arr = np.linspace(float(lo), float(hi), n_frames)
1354
+ scale = np.maximum(1.0, np.abs(times_arr))
1355
+ integer = bool(
1356
+ np.all(np.abs(times_arr - np.round(times_arr)) <= 1e-6 * scale)
1357
+ )
1358
+ if integer:
1359
+ times = [int(round(float(t))) for t in times_arr]
1360
+ else:
1361
+ times = [float(t) for t in times_arr]
1362
+ params = [
1363
+ {"name": name, "lo": float(bounds[0]), "hi": float(bounds[1])}
1364
+ for name, bounds in transition.ranges.items()
1365
+ ]
1366
+ return {
1367
+ "type": "time",
1368
+ "column": first,
1369
+ "nFrames": n_frames,
1370
+ "times": times,
1371
+ "integer": integer,
1372
+ "ease": "linear",
1373
+ "duration": 12,
1374
+ "params": params,
1375
+ }
1376
+
1377
+
1378
+ def _slider_meta(slider) -> dict:
1379
+ """Independent ranges. The last keyword varies fastest (stride 1)."""
1380
+ names = list(slider.ranges)
1381
+ steps = int(slider.steps)
1382
+ params: list[dict] = []
1383
+ stride = 1
1384
+ for name in reversed(names):
1385
+ lo, hi = slider.ranges[name]
1386
+ params.append(
1387
+ {
1388
+ "name": name,
1389
+ "lo": float(lo),
1390
+ "hi": float(hi),
1391
+ "n": steps,
1392
+ "stride": stride,
1393
+ }
1394
+ )
1395
+ stride *= steps
1396
+ params.reverse()
1397
+ return {"nFrames": int(stride), "params": params}
1398
+
1399
+
1400
+ def _legend_position_spec(value):
1401
+ """JSON form of theme(legend_position=). The static default is outside right."""
1402
+ if value is None:
1403
+ return "right"
1404
+ if isinstance(value, tuple):
1405
+ return [float(value[0]), float(value[1])]
1406
+ return value
1407
+
1408
+
1409
+
1410
+ def _forced_levels(g, ccats) -> list:
1411
+ """Colour levels shared by every facet panel, so a group keeps its colour."""
1412
+ forced = getattr(g, "_force_color_levels", None)
1413
+ if not forced:
1414
+ return list(ccats)
1415
+ return list(dict.fromkeys([str(level) for level in forced] + list(ccats)))
1416
+
1417
+
1418
+ def _limits_in_scale(pscale, sc) -> tuple[float | None, float | None]:
1419
+ """Scale limits in the scale's own units (log10, epoch seconds)."""
1420
+ lo, hi = pscale.limits
1421
+ def conv(value):
1422
+ if value is None:
1423
+ return None
1424
+ if pscale.kind == "date" or (sc is not None and sc.kind == "dt"):
1425
+ return float(pd.Timestamp(value).timestamp())
1426
+ number = float(value)
1427
+ if sc is not None and getattr(sc, "trans", None) == "log10":
1428
+ if number <= 0:
1429
+ raise ValueError("log scale limits must be positive")
1430
+ return math.log10(number)
1431
+ return number
1432
+ a, b = conv(lo), conv(hi)
1433
+ if a is not None and b is not None and b < a:
1434
+ a, b = b, a
1435
+ return a, b
1436
+
1437
+
1438
+ class _Limits:
1439
+ """A pair of limits shaped like a position scale, for _limits_in_scale."""
1440
+
1441
+ kind = "continuous"
1442
+
1443
+ def __init__(self, limits):
1444
+ self.limits = tuple(limits)
1445
+
1446
+
1447
+ def _rug_series(rug, g):
1448
+ """(axis, values) for each side a rug draws."""
1449
+ from plot3.table import as_table, materialize_columns
1450
+
1451
+ mapping = _apply_fill(dict(g.mapping), "point")
1452
+ mapping.update(_apply_fill(dict(rug.mapping), "point"))
1453
+ data = getattr(rug, "layer_data", None)
1454
+ data = g.data if data is None else as_table(data)
1455
+ if data is None:
1456
+ return []
1457
+ out = []
1458
+ for axis in dict.fromkeys("x" if side in "bt" else "y" for side in rug.sides):
1459
+ column = mapping.get(axis)
1460
+ if not column:
1461
+ continue
1462
+ frame = materialize_columns(data, [column])
1463
+ out.append((axis, frame[column]))
1464
+ return out
1465
+
1466
+
1467
+ def _rug_positions(scale, series) -> np.ndarray:
1468
+ """Values in the scale's own units; NaN where they have no place."""
1469
+ kind, pos, cats = col_values(series)
1470
+ if scale.kind == "cat":
1471
+ index = {c: i for i, c in enumerate(scale.cats)}
1472
+ if kind == "cat":
1473
+ lookup = np.array([index.get(c, np.nan) for c in cats] + [np.nan], dtype=np.float64)
1474
+ codes = np.where(np.isfinite(pos), pos, len(cats)).astype(np.int64)
1475
+ return lookup[codes]
1476
+ return np.array([index.get(str(v), np.nan) for v in series], dtype=np.float64)
1477
+ if kind == "cat":
1478
+ return np.full(len(pos), np.nan)
1479
+ if getattr(scale, "trans", None) == "log10":
1480
+ with np.errstate(divide="ignore", invalid="ignore"):
1481
+ return np.where(pos > 0, np.log10(pos), np.nan)
1482
+ return pos
1483
+
1484
+
1485
+ def _rug_specs(rug_layers, g, scales, theme, cspec, cat_colours) -> list[dict]:
1486
+ """geom_rug ticks per side, in scale units, with a colour per tick when
1487
+ the rug maps a colour that the figure's legend already shows."""
1488
+ from plot3.table import as_table, materialize_columns
1489
+
1490
+ out = []
1491
+ for rug in rug_layers:
1492
+ mapping = _apply_fill(dict(g.mapping), "point")
1493
+ mapping.update(_apply_fill(dict(rug.mapping), "point"))
1494
+ colour_of = None
1495
+ column = mapping.get("color")
1496
+ if not rug.const_color and column and cspec and cspec.get("kind") == "cat":
1497
+ data = getattr(rug, "layer_data", None)
1498
+ data = g.data if data is None else as_table(data)
1499
+ levels = materialize_columns(data, [column])[column]
1500
+ index = {str(c): i for i, c in enumerate(cspec.get("cats") or [])}
1501
+ colour_of = [
1502
+ cat_colours[index[str(v)] % len(cat_colours)] if str(v) in index else None
1503
+ for v in levels
1504
+ ]
1505
+ base = _hex_or_none(rug.const_color) or theme["ink"]
1506
+ series = dict(_rug_series(rug, g))
1507
+ for side in rug.sides:
1508
+ axis = "x" if side in "bt" else "y"
1509
+ if axis not in series or axis not in scales:
1510
+ continue
1511
+ values = _rug_positions(scales[axis], series[axis])
1512
+ keep = np.isfinite(values)
1513
+ item = {
1514
+ "side": side,
1515
+ "values": [float(v) for v in values[keep]],
1516
+ "color": base,
1517
+ "length": float(rug.length),
1518
+ "width": float(rug.linewidth) * 1.5,
1519
+ "alpha": 1.0 if rug.alpha is None else float(rug.alpha),
1520
+ }
1521
+ if colour_of is not None:
1522
+ item["colors"] = [colour_of[i] or base for i in np.flatnonzero(keep)]
1523
+ out.append(item)
1524
+ return out
1525
+
1526
+
1527
+ def _apply_guides(spec: dict, hidden: dict) -> None:
1528
+ """guides(colour="none") and friends: drop those legends from the spec."""
1529
+ if hidden.get("color"):
1530
+ spec["legend"] = None
1531
+ if spec.get("color") and spec["color"].get("kind") == "num":
1532
+ spec["color"]["guide"] = False
1533
+ if hidden.get("size"):
1534
+ spec["sizeLegend"] = None
1535
+ if hidden.get("shape"):
1536
+ spec["shapeLegend"] = None
1537
+ if hidden.get("linetype"):
1538
+ spec["linetypeLegend"] = None
1539
+
1540
+
1541
+ def _apply_guide_options(spec: dict, options: dict) -> None:
1542
+ """guide_legend(title=, reverse=) and guide_colourbar(title=)."""
1543
+ for name, opts in options.items():
1544
+ title, reverse = opts.get("title"), opts.get("reverse")
1545
+ if name == "color":
1546
+ if title is not None:
1547
+ key = "colorBar" if (spec.get("labs") or {}).get("colorBar") is not None else "color"
1548
+ spec["labs"][key] = str(title)
1549
+ if reverse and spec.get("legend"):
1550
+ spec["legend"] = list(reversed(spec["legend"]))
1551
+ continue
1552
+ legend = spec.get({"size": "sizeLegend", "shape": "shapeLegend",
1553
+ "linetype": "linetypeLegend"}[name])
1554
+ if not legend:
1555
+ continue
1556
+ if title is not None:
1557
+ legend["label"] = str(title)
1558
+ rows = "breaks" if name == "size" else "entries"
1559
+ if reverse and legend.get(rows):
1560
+ legend[rows] = list(reversed(legend[rows]))
1561
+
1562
+
1563
+ def _hex_or_none(colour):
1564
+ """#rrggbb for the viewer, which does not know R's 'grey50'."""
1565
+ if not colour:
1566
+ return colour
1567
+ from plot3.scaling import to_hex
1568
+
1569
+ try:
1570
+ return to_hex(colour)
1571
+ except ValueError:
1572
+ return colour
1573
+
1574
+
1575
+ def _label_text(value) -> str:
1576
+ """geom_text label: numbers to 4 significant figures, text as given."""
1577
+ if value is None or (isinstance(value, float) and math.isnan(value)):
1578
+ return ""
1579
+ if isinstance(value, (float, np.floating)):
1580
+ number = float(value)
1581
+ if abs(number) >= 1e4:
1582
+ return f"{number:,.0f}"
1583
+ return f"{number:.4g}".replace("-", "−")
1584
+ return str(value)
1585
+
1586
+
1587
+ def _arrow_specs(geom, vals, order, spec_l, cat_colours, theme) -> list[dict]:
1588
+ """Arrowheads at the ends of each line group, in scale units.
1589
+
1590
+ The renderers draw them in screen space, so a head keeps its shape
1591
+ whatever the axes' aspect ratio.
1592
+ """
1593
+ style = geom.arrow.spec()
1594
+ xs = np.asarray(vals["x"], dtype=np.float64)[order]
1595
+ ys = np.asarray(vals["y"], dtype=np.float64)[order]
1596
+ codes = None
1597
+ if vals.get("color") and vals["color"][0] == "cat" and cat_colours:
1598
+ codes = np.asarray(vals["color"][1], dtype=np.float64)[order]
1599
+ base = spec_l.get("constColor") or theme["cat"][0]
1600
+ out = []
1601
+ for start, count in spec_l.get("groups") or []:
1602
+ start, count = int(start), int(count)
1603
+ if count < 2:
1604
+ continue
1605
+ colour = base
1606
+ if codes is not None and math.isfinite(codes[start]):
1607
+ colour = cat_colours[int(codes[start]) % len(cat_colours)]
1608
+ ends = []
1609
+ if style["ends"] in {"last", "both"}:
1610
+ ends.append((start + count - 2, start + count - 1))
1611
+ if style["ends"] in {"first", "both"}:
1612
+ ends.append((start + 1, start))
1613
+ for a, b in ends:
1614
+ if all(math.isfinite(v) for v in (xs[a], ys[a], xs[b], ys[b])):
1615
+ out.append({
1616
+ "x0": float(xs[a]), "y0": float(ys[a]), "x1": float(xs[b]), "y1": float(ys[b]),
1617
+ "color": colour, "width": float(spec_l.get("linewidth") or 1.0),
1618
+ **{k: style[k] for k in ("angle", "length", "type")},
1619
+ })
1620
+ return out
1621
+
1622
+
1623
+ def _layer_name(geom) -> str:
1624
+ """The geom as the user wrote it, for messages: geom_point, geom_line."""
1625
+ name = type(geom).__name__
1626
+ return name if name.startswith(("geom_", "stat_")) else f"geom_{geom.kind}"
1627
+
1628
+
1629
+ def _figure_theme(g) -> dict:
1630
+ """The named theme with theme(element_*()) colours laid over it."""
1631
+ theme = dict(THEMES[g.theme_name])
1632
+ tokens = (getattr(g, "theme_options", None) or {}).get("tokens") or {}
1633
+ for key, value in tokens.items():
1634
+ if key == "panel" and value is None:
1635
+ theme["panel"] = theme["surface"]
1636
+ else:
1637
+ theme[key] = value
1638
+ theme.setdefault("panel", theme["surface"])
1639
+ return theme
1640
+
1641
+
1642
+ def _theme_opts(g) -> dict | None:
1643
+ """theme() settings both renderers read (grid, label angle, title hjust)."""
1644
+ options = getattr(g, "theme_options", None) or {}
1645
+ out = {}
1646
+ if "panel_grid" in options:
1647
+ out["panelGrid"] = bool(options["panel_grid"])
1648
+ if "axis_text_x_angle" in options:
1649
+ out["xAngle"] = float(options["axis_text_x_angle"])
1650
+ if "plot_title_hjust" in options:
1651
+ out["titleHjust"] = float(options["plot_title_hjust"])
1652
+ # theme(axis_text_x=element_blank()): no tick labels on that axis.
1653
+ for axis in ("x", "y"):
1654
+ if options.get(f"axis_text_{axis}") is False:
1655
+ out[f"{axis}Text"] = False
1656
+ return out or None
1657
+
1658
+
1659
+ def _value_part(full: str, base: str, sep: str) -> str | None:
1660
+ """``a = 2, b = 5`` from a caption ``<base><sep>(a = 2, b = 5)``."""
1661
+ head = f"{base}{sep}("
1662
+ if base and full.startswith(head) and full.endswith(")"):
1663
+ return full[len(head):-1]
1664
+ return None
1665
+
1666
+
1667
+ def _formula_legend(legend, formula_geoms, has_title: bool):
1668
+ """Tidy the legend rows that function curves add.
1669
+
1670
+ * One curve under a title of your own: the title already names it, so
1671
+ the curve gets no row (ggplot2 draws no legend for an unmapped layer).
1672
+ * Several curves of one formula (a loop over ``a``): the formula moves
1673
+ to the legend title and each row lists only its values.
1674
+
1675
+ Returns ``(legend, (pretty, latex) | None)``.
1676
+ """
1677
+ if not legend:
1678
+ return legend, None
1679
+ primary = [
1680
+ entry for entry in legend
1681
+ if entry.get("formula")
1682
+ and getattr(entry.get("_geom"), "_formula_primary", False)
1683
+ and not getattr(entry.get("_geom"), "_legend_math", None)
1684
+ ]
1685
+ if not primary:
1686
+ return legend, None
1687
+ if has_title and len(formula_geoms) == 1:
1688
+ kept = [entry for entry in legend if not any(entry is p for p in primary)]
1689
+ return kept or None, None
1690
+ if len(primary) < 2 or len(primary) != len(formula_geoms):
1691
+ return legend, None
1692
+ geoms = [entry["_geom"] for entry in primary]
1693
+ bases = {str(getattr(geom, "_title_label", "") or "") for geom in geoms}
1694
+ if len(bases) != 1:
1695
+ return legend, None
1696
+ base = bases.pop()
1697
+ base_latex = str(getattr(geoms[0], "_title_latex", "") or "")
1698
+ rows = []
1699
+ for geom in geoms:
1700
+ pretty = _value_part(str(getattr(geom, "_tip_pretty", "") or ""), base, " ")
1701
+ latex = _value_part(
1702
+ str(getattr(geom, "_tip_latex", "") or ""), base_latex, " \\quad "
1703
+ )
1704
+ if pretty is None:
1705
+ return legend, None
1706
+ rows.append((pretty, latex))
1707
+ for entry, (pretty, latex) in zip(primary, rows):
1708
+ entry["label"] = pretty
1709
+ entry.pop("math", None)
1710
+ if latex:
1711
+ entry["latex"] = latex
1712
+ else:
1713
+ entry.pop("latex", None)
1714
+ return legend, (base, base_latex)
1715
+
1716
+
1717
+ def _legend_title_label(current: str, legend_title, labs_math: dict) -> str:
1718
+ """Use a shared formula as the legend title unless labs(colour=) is set."""
1719
+ if current or not legend_title:
1720
+ return current
1721
+ pretty, latex = legend_title
1722
+ if latex:
1723
+ labs_math["color"] = [{"text": pretty, "latex": latex}]
1724
+ return pretty
1725
+
1726
+
1727
+ def _discrete_codes(series: pd.Series) -> tuple[np.ndarray, list[str]]:
1728
+ """Category codes and names for shape / linetype (numbers become levels)."""
1729
+ kind, values, cats = col_values(series)
1730
+ if kind == "cat":
1731
+ return np.asarray(values, dtype=np.float64), list(cats)
1732
+ labels = series.map(lambda v: "NA" if pd.isna(v) else str(v))
1733
+ order = sorted(set(labels), key=lambda t: (_num_key(t), t))
1734
+ index = {name: i for i, name in enumerate(order)}
1735
+ return labels.map(index).to_numpy(np.float64), order
1736
+
1737
+
1738
+ def _num_key(text: str) -> float:
1739
+ try:
1740
+ return float(text)
1741
+ except ValueError:
1742
+ return math.inf
1743
+
1744
+
1745
+ def _linetype_names(cats, key_scale=None) -> list:
1746
+ from plot3.geoms import LINETYPE_ORDER
1747
+
1748
+ if key_scale is not None:
1749
+ return [key_scale.value_for(i, str(c)) for i, c in enumerate(cats)]
1750
+ return [LINETYPE_ORDER[i % len(LINETYPE_ORDER)] for i in range(len(cats))]
1751
+
1752
+
1753
+ def _shape_names(cats, key_scale=None) -> list:
1754
+ from plot3.geoms import SHAPE_ORDER
1755
+
1756
+ if key_scale is not None:
1757
+ return [key_scale.value_for(i, str(c)) for i, c in enumerate(cats)]
1758
+ if len(cats) > len(SHAPE_ORDER):
1759
+ raise ValueError(
1760
+ f"aes(shape=) has {len(cats)} levels; at most {len(SHAPE_ORDER)} "
1761
+ "shapes are distinct. Map a column with fewer levels to shape, or "
1762
+ "give scale_shape_manual(values=[...]) one shape per level"
1763
+ )
1764
+ return [SHAPE_ORDER[i] for i in range(len(cats))]
1765
+
1766
+
1767
+ def _encode_dashes(spec_l: dict, geom, vals: dict, order: np.ndarray, key_scale=None) -> None:
1768
+ """Dash patterns for a line layer: one per group, or one for the layer."""
1769
+ from plot3.geoms import dash_pattern
1770
+
1771
+ mapped = vals.get("linetype")
1772
+ if mapped is not None:
1773
+ codes = np.asarray(mapped[0], dtype=np.float64)[order]
1774
+ names = _linetype_names(mapped[1], key_scale)
1775
+ dashes = []
1776
+ for start, _count in spec_l.get("groups") or []:
1777
+ code = codes[int(start)]
1778
+ name = names[int(code) % len(names)] if math.isfinite(code) else "solid"
1779
+ pattern = dash_pattern(name)
1780
+ dashes.append(list(pattern) if pattern else None)
1781
+ if any(dashes):
1782
+ spec_l["dashes"] = dashes
1783
+ return
1784
+ pattern = dash_pattern(getattr(geom, "linetype", None))
1785
+ if pattern:
1786
+ spec_l["dash"] = list(pattern)
1787
+
1788
+
1789
+ def _encode_shapes(spec_l: dict, geom, vals: dict, order: np.ndarray, li: int, payloads: list, compress: bool, key_scale=None) -> None:
1790
+ """Point symbols: per-point codes with their names, or one shape."""
1791
+ mapped = vals.get("shape")
1792
+ if mapped is not None:
1793
+ codes, cats = mapped
1794
+ names = _shape_names(cats, key_scale)
1795
+ ordered = np.asarray(codes, dtype=np.float64)[order]
1796
+ packed = np.where(np.isfinite(ordered), ordered, 0).astype("<u2")
1797
+ pid = f"p{li}sh"
1798
+ payloads.append((pid, pack_u16(packed, compress)))
1799
+ spec_l["shape"] = {"id": pid, "dtype": "u16", "names": names}
1800
+ elif getattr(geom, "shape", None):
1801
+ spec_l["shape"] = str(geom.shape)
1802
+
1803
+
1804
+ def _aux_legends(resolved, layer_vals, legend, shape_scale=None, linetype_scale=None):
1805
+ """Shape and linetype keys: on the colour legend when they map the same
1806
+ column, otherwise a legend of their own."""
1807
+ from plot3.geoms import dash_pattern
1808
+
1809
+ extra = {}
1810
+ for (geom, m), vals in zip(resolved, layer_vals):
1811
+ for aes_name, key in (("shape", "shape"), ("linetype", "dash")):
1812
+ mapped = vals.get(aes_name)
1813
+ if mapped is None or aes_name in extra:
1814
+ continue
1815
+ cats = mapped[1]
1816
+ if key == "dash":
1817
+ styles = [list(dash_pattern(n) or []) for n in _linetype_names(cats, linetype_scale)]
1818
+ else:
1819
+ styles = _shape_names(cats, shape_scale)
1820
+ column = m.get(aes_name)
1821
+ key_scale = shape_scale if key == "shape" else linetype_scale
1822
+ title = key_scale.name if key_scale is not None and key_scale.name else column
1823
+ by_label = dict(zip(cats, styles))
1824
+ if legend and column == m.get("color") and all(e["label"] in by_label for e in legend):
1825
+ for entry in legend:
1826
+ entry[key] = by_label[entry["label"]]
1827
+ extra[aes_name] = None
1828
+ else:
1829
+ extra[aes_name] = {
1830
+ "label": str(title),
1831
+ "entries": [{"label": c, key: st} for c, st in zip(cats, styles)],
1832
+ }
1833
+ # One constant symbol for every point layer (geom_point(shape="triangle"))
1834
+ # also belongs on the colour keys.
1835
+ if legend and "shape" not in extra:
1836
+ constant = {
1837
+ getattr(geom, "shape", None)
1838
+ for geom, _m in resolved
1839
+ if geom.kind == "point" and getattr(geom, "shape", None)
1840
+ }
1841
+ if len(constant) == 1 and all(not e.get("shape") for e in legend):
1842
+ shape = constant.pop()
1843
+ for entry in legend:
1844
+ entry["shape"] = shape
1845
+ return legend, extra.get("shape"), extra.get("linetype")
1846
+
1847
+
1848
+ def build_spec(g: ggplot) -> tuple[dict, list[tuple[str, str]]]:
1849
+ if not g.layers:
1850
+ raise ValueError("add a geom: ggplot(df, aes(...)) + geom_point()")
1851
+ # aes(ymin="mean - se"), aes(colour="factor(cyl)"): computed columns.
1852
+ from plot3.aesexpr import add_expression_columns
1853
+
1854
+ g = add_expression_columns(g)
1855
+ # Function layers and vector fields sample their own grid, so a figure
1856
+ # may have no data frame.
1857
+ needs_data = any(
1858
+ getattr(layer, "kind", None) not in {"function", "vector", "hline", "vline", "abline"}
1859
+ and getattr(layer, "layer_data", None) is None
1860
+ and not getattr(layer, "_annotation", False)
1861
+ for layer in g.layers
1862
+ )
1863
+ if g.data is None and needs_data:
1864
+ raise ValueError(
1865
+ "ggplot has no data; use ggplot(df, aes(...)) or "
1866
+ "pipe data with `data >> ggplot(aes(...))`"
1867
+ )
1868
+
1869
+ data = g.data
1870
+ coord = getattr(g, "coord", None)
1871
+ mesh_kinds = {"surface", "isosurface"}
1872
+ has_mesh = any(
1873
+ getattr(layer, "kind", None) in mesh_kinds for layer in g.layers
1874
+ )
1875
+ if (
1876
+ data is not None
1877
+ and coord is not None
1878
+ and getattr(coord, "max_points", None)
1879
+ and not has_mesh
1880
+ ):
1881
+ nrows = n_rows(data)
1882
+ cap = int(coord.max_points)
1883
+ if nrows > cap:
1884
+ step = max(1, (nrows + cap - 1) // cap)
1885
+ data = subsample_rows(data, step)
1886
+
1887
+ theme = _figure_theme(g)
1888
+ # Apply optional stat_density_3d options onto isosurface layers.
1889
+ density_stat = getattr(g, "stat_density_3d", None)
1890
+ layers_in = []
1891
+ ref_layers = []
1892
+ rug_layers = []
1893
+ for geom in g.layers:
1894
+ if getattr(geom, "kind", None) in _REF_KINDS:
1895
+ ref_layers.append(geom) # drawn across the panel, not from rows
1896
+ continue
1897
+ if getattr(geom, "kind", None) == "rug":
1898
+ rug_layers.append(geom) # ticks on the panel edges
1899
+ continue
1900
+ if getattr(geom, "kind", None) == "isosurface" and density_stat is not None:
1901
+ geom = copy_geom_with_density_n(geom, density_stat.n)
1902
+ layers_in.append(geom)
1903
+ if not layers_in:
1904
+ raise ValueError(
1905
+ "geom_hline(), geom_vline(), geom_abline(), and geom_rug() draw on a "
1906
+ "plot: add a data layer such as geom_point()"
1907
+ )
1908
+ transition = getattr(g, "transition", None)
1909
+ slider = getattr(g, "slider", None)
1910
+ if slider is not None and transition is not None:
1911
+ raise ValueError(
1912
+ "slider() cannot be combined with transition_time() "
1913
+ "or transition_states()"
1914
+ )
1915
+ domains = None
1916
+ if any(getattr(layer, "kind", None) == "function" for layer in layers_in):
1917
+ from plot3.function import data_domains
1918
+
1919
+ domains = data_domains(g, data)
1920
+ addon_map: dict[int, list] = {}
1921
+ for index, addon in getattr(g, "_addons", None) or []:
1922
+ addon_map.setdefault(int(index), []).append(addon)
1923
+ expanded = []
1924
+ for index, geom in enumerate(layers_in):
1925
+ layer_data = getattr(geom, "layer_data", None)
1926
+ if getattr(geom, "kind", None) == "box3d":
1927
+ # Box classes take the figure theme's group colours.
1928
+ geom = copy.copy(geom)
1929
+ geom._palette = list(theme["cat"])
1930
+ result = expand_stat_geom(
1931
+ geom,
1932
+ g.mapping,
1933
+ data if layer_data is None else layer_data,
1934
+ domains,
1935
+ transition,
1936
+ slider,
1937
+ coord,
1938
+ addon_map.get(index),
1939
+ )
1940
+ if isinstance(result, list):
1941
+ expanded.extend(result)
1942
+ else:
1943
+ expanded.append(result)
1944
+ from plot3.calculus import mark_intersections
1945
+
1946
+ expanded = mark_intersections(expanded)
1947
+ from plot3.geoms import coord_flip as _coord_flip
1948
+
1949
+ if isinstance(coord, _coord_flip):
1950
+ # Stats ran upright; now every layer exchanges x and y.
1951
+ from plot3.flip import flip_layers, flip_refs, flipped_figure
1952
+
1953
+ expanded = flip_layers(expanded, g)
1954
+ g = flipped_figure(g)
1955
+ ref_layers = flip_refs(ref_layers)
1956
+ from plot3.flip import flip_rugs
1957
+
1958
+ rug_layers = flip_rugs(rug_layers)
1959
+ coord = None
1960
+ resolved = [] # per layer: (geom, mapping)
1961
+ for geom in expanded:
1962
+ # Function layers carry their own columns; don't inherit colour/group.
1963
+ if getattr(geom, "_replace_mapping", False):
1964
+ m = dict(geom.mapping)
1965
+ else:
1966
+ m = _apply_fill(dict(g.mapping), geom.kind)
1967
+ m.update(_apply_fill(dict(geom.mapping), geom.kind))
1968
+ if "x" not in m or "y" not in m:
1969
+ raise ValueError("aes(x=, y=) are required (bar/histogram/density supply y)")
1970
+ if geom.kind == "surface" and "z" not in m:
1971
+ raise ValueError("geom_surface() requires aes(z=)")
1972
+ resolved.append((geom, m))
1973
+
1974
+ is3d = any("z" in m for _, m in resolved)
1975
+ if is3d and not all("z" in m for _, m in resolved):
1976
+ raise ValueError("mix of 2D and 3D layers: every layer needs aes(z=)")
1977
+ if is3d and any(geom.kind in _2D_ONLY_KINDS for geom, _ in resolved):
1978
+ bad = sorted(
1979
+ {geom.kind for geom, _ in resolved if geom.kind in _2D_ONLY_KINDS}
1980
+ )
1981
+ raise ValueError(
1982
+ f"geom kind(s) {bad} are 2D-only; use geom_point / geom_point3d "
1983
+ "(or geom_line/path/surface) with aes(z=...) for 3D"
1984
+ )
1985
+ if is3d and getattr(g, "facet", None) is not None:
1986
+ raise ValueError("facet_wrap() is not supported with 3D figures yet")
1987
+ if is3d and rug_layers:
1988
+ raise ValueError("geom_rug() is for 2D figures")
1989
+ # A point cloud with no colour of its own is coloured by height, as
1990
+ # lidar viewers draw it: aes(colour=) or colour="steelblue" turns it off.
1991
+ height_title = None
1992
+ if is3d and not any("color" in m for _g, m in resolved):
1993
+ for geom, m in resolved:
1994
+ if (
1995
+ geom.kind == "point"
1996
+ and getattr(geom, "const_color", None) is None
1997
+ and not getattr(geom, "_replace_mapping", False)
1998
+ ):
1999
+ m["color"] = m["z"]
2000
+ height_title = str(m["z"])
2001
+
2002
+ if transition is not None:
2003
+ tname = (
2004
+ "transition_states" if transition.kind == "states" else "transition_time"
2005
+ )
2006
+ # geom_function has already become a line, path, or surface. A
2007
+ # parameter sweep is recognised by the formula stamp, not by kind.
2008
+ if getattr(transition, "ranges", None):
2009
+ formula_layers = [
2010
+ geom
2011
+ for geom, _mapped in resolved
2012
+ if getattr(geom, "_is_formula", False)
2013
+ ]
2014
+ if not formula_layers:
2015
+ raise ValueError(
2016
+ "transition_time() parameter ranges need a geom_function layer"
2017
+ )
2018
+ extra = sorted(
2019
+ {
2020
+ geom.kind
2021
+ for geom, _mapped in resolved
2022
+ if not getattr(geom, "_is_formula", False)
2023
+ }
2024
+ )
2025
+ if extra:
2026
+ raise ValueError(
2027
+ "transition_time() parameter ranges animate geom_function "
2028
+ f"layers only (this figure also has {extra})"
2029
+ )
2030
+ else:
2031
+ kinds = {geom.kind for geom, _ in resolved}
2032
+ if "point" not in kinds:
2033
+ raise ValueError(f"{tname}() needs a geom_point layer")
2034
+ extra = sorted(kinds - {"point"})
2035
+ if extra:
2036
+ raise ValueError(
2037
+ f"{tname}() animates geom_point layers only "
2038
+ f"(this figure also has {extra})"
2039
+ )
2040
+
2041
+ if slider is not None and getattr(slider, "ranges", None):
2042
+ formula_layers = [
2043
+ geom
2044
+ for geom, _mapped in resolved
2045
+ if getattr(geom, "_is_formula", False)
2046
+ ]
2047
+ if not formula_layers:
2048
+ raise ValueError("slider() needs a geom_function layer")
2049
+ extra = sorted(
2050
+ {
2051
+ geom.kind
2052
+ for geom, _mapped in resolved
2053
+ if not getattr(geom, "_is_formula", False)
2054
+ }
2055
+ )
2056
+ if extra:
2057
+ raise ValueError(
2058
+ "slider() drives geom_function layers only "
2059
+ f"(this figure also has {extra})"
2060
+ )
2061
+ if not any(getattr(geom, "_anim", None) for geom, _mapped in resolved):
2062
+ raise ValueError(
2063
+ "slider() parameters are not coefficients of a geom_function. "
2064
+ 'Leave them unbound: geom_function("y = a x^2") + slider(a=(0, 3))'
2065
+ )
2066
+
2067
+ axes = ["x", "y", "z"] if is3d else ["x", "y"]
2068
+ scales: dict[str, Scale] = {}
2069
+ # scale_x_discrete(limits=[...]) / xlim("a", "b"): that order, first.
2070
+ for _axis_name, _pscale in (("x", getattr(g, "xscale", None)), ("y", getattr(g, "yscale", None))):
2071
+ if _pscale is not None and _pscale.kind == "discrete" and _pscale.limits:
2072
+ scales[_axis_name] = Scale("cat")
2073
+ scales[_axis_name].cats = list(_pscale.limits)
2074
+ color_scale = None # ("num", lo, hi) | ("cat", cats)
2075
+ num_color_vals: list[np.ndarray] = []
2076
+ dropped_log = 0
2077
+ missing_notes: list[str] = []
2078
+ size_label = str(g.labs.get("size") or "")
2079
+ alpha_label = str(g.labs.get("alpha") or "")
2080
+ transition_meta: dict | None = None
2081
+ if transition is not None and getattr(transition, "ranges", None):
2082
+ transition_meta = _range_transition_meta(transition)
2083
+ slider_meta = (
2084
+ _slider_meta(slider)
2085
+ if slider is not None and getattr(slider, "ranges", None)
2086
+ else None
2087
+ )
2088
+
2089
+ def _axis_trans(axis: str) -> str | None:
2090
+ if axis == "x" and getattr(g, "scale_x", None) is not None:
2091
+ return "log10"
2092
+ if axis == "y" and getattr(g, "scale_y", None) is not None:
2093
+ return "log10"
2094
+ return None
2095
+
2096
+ def _absorb_level_positions(v: np.ndarray, levels: list[str], axis: str = "x") -> np.ndarray:
2097
+ if _axis_trans(axis) == "log10":
2098
+ raise ValueError(f"scale_{axis}_log10() cannot be used with categorical {axis}")
2099
+ sc = scales.get(axis)
2100
+ if sc is None or (sc.kind == "num" and not math.isfinite(sc.lo)):
2101
+ sc = scales[axis] = Scale("cat")
2102
+ elif sc.kind != "cat":
2103
+ raise ValueError(
2104
+ "aes x: layers disagree on scale type (cat vs num)"
2105
+ )
2106
+ merged = list(dict.fromkeys(sc.cats + [str(level) for level in levels]))
2107
+ sc.cats = merged
2108
+ where = {str(level): merged.index(str(level)) for level in levels}
2109
+ flat = np.asarray(v, dtype=np.float64)
2110
+ # Nearest category, clamped: a tile's outer edge sits at n - 0.5,
2111
+ # which rounds past the last level. Keep the offset from it.
2112
+ top = max(len(levels) - 1, 0)
2113
+ out = np.full(flat.shape, np.nan)
2114
+ for i, p in enumerate(flat):
2115
+ if not math.isfinite(p) or not levels:
2116
+ continue
2117
+ k = int(min(max(math.floor(p + 0.5), 0), top))
2118
+ out[i] = where[str(levels[k])] + (p - k)
2119
+ return out
2120
+
2121
+ def _absorb_position(axis: str, kind: str, v: np.ndarray, cats: list[str]):
2122
+ nonlocal dropped_log
2123
+ trans = _axis_trans(axis)
2124
+ sc = scales.get(axis)
2125
+ if trans == "log10" and kind != "num":
2126
+ raise ValueError(
2127
+ f"scale_{axis}_log10() needs a numeric {axis} column"
2128
+ )
2129
+ if sc is None:
2130
+ sc = scales[axis] = Scale(kind, trans=trans if kind == "num" else None)
2131
+ elif sc.kind != kind:
2132
+ raise ValueError(
2133
+ f"aes {axis}: layers disagree on scale type "
2134
+ f"({sc.kind} vs {kind})"
2135
+ )
2136
+ if kind == "cat":
2137
+ merged = list(dict.fromkeys(sc.cats + cats))
2138
+ remap = {
2139
+ cats.index(c) if c in cats else None: i
2140
+ for i, c in enumerate(merged)
2141
+ if c in cats
2142
+ }
2143
+ out = np.full(np.size(v), np.nan, dtype=np.float64)
2144
+ flat = np.asarray(v, dtype=np.float64).ravel()
2145
+ for i, c in enumerate(flat):
2146
+ if not math.isfinite(c):
2147
+ continue
2148
+ out[i] = remap.get(int(c), -1)
2149
+ v = out.reshape(np.shape(v))
2150
+ sc.cats = merged
2151
+ else:
2152
+ if trans == "log10":
2153
+ v, n_bad = _log10_values(v)
2154
+ dropped_log += n_bad
2155
+ sc.widen(v)
2156
+ return v
2157
+
2158
+ def _take_color(column_values_kind, cv, ccats, *, into):
2159
+ """Record a colour channel on ``into`` and widen the shared colour scale."""
2160
+ nonlocal color_scale
2161
+ kind = column_values_kind
2162
+ if kind == "cat" or (kind == "num" and ccats):
2163
+ if color_scale is None:
2164
+ color_scale = ["cat", _forced_levels(g, ccats)]
2165
+ else:
2166
+ if color_scale[0] != "cat":
2167
+ raise ValueError("layers disagree on colour scale type")
2168
+ color_scale[1] = list(dict.fromkeys(color_scale[1] + ccats))
2169
+ into["color"] = ("cat", cv, ccats)
2170
+ else:
2171
+ finite = np.asarray(cv, dtype=np.float64)
2172
+ finite = finite[np.isfinite(finite)]
2173
+ if color_scale is None:
2174
+ color_scale = ["num", math.inf, -math.inf]
2175
+ elif color_scale[0] != "num":
2176
+ raise ValueError("layers disagree on colour scale type")
2177
+ if finite.size:
2178
+ color_scale[1] = min(color_scale[1], float(finite.min()))
2179
+ color_scale[2] = max(color_scale[2], float(finite.max()))
2180
+ num_color_vals.append(np.asarray(cv, dtype=np.float64))
2181
+ into["color"] = ("num", cv, None)
2182
+
2183
+ def _limit_rows(vals: dict, geom) -> None:
2184
+ """Rows outside scale limits are dropped, as in ggplot2, and reported."""
2185
+ if "frames" in vals or getattr(geom, "data_override", None) is not None:
2186
+ return
2187
+ n = len(vals["x"])
2188
+ ok = np.ones(n, dtype=bool)
2189
+ for axis_name in ("x", "y"):
2190
+ pscale = getattr(g, f"{axis_name}scale", None)
2191
+ if pscale is None or axis_name not in vals or pscale.limits is None:
2192
+ continue
2193
+ v = np.asarray(vals[axis_name], dtype=np.float64)
2194
+ if pscale.kind == "discrete":
2195
+ ok &= np.isfinite(v) & (v < len(pscale.limits))
2196
+ continue
2197
+ lo, hi = _limits_in_scale(pscale, scales.get(axis_name))
2198
+ tol = 1e-9 * max(1.0, abs(lo or 0.0), abs(hi or 0.0))
2199
+ if lo is not None:
2200
+ ok &= v >= lo - tol
2201
+ if hi is not None:
2202
+ ok &= v <= hi + tol
2203
+ if ok.all():
2204
+ return
2205
+ removed = int((~ok).sum())
2206
+ word = "row" if removed == 1 else "rows"
2207
+ missing_notes.append(
2208
+ f"Removed {removed} {word} outside the scale limits ({_layer_name(geom)})"
2209
+ )
2210
+ _filter_rows(vals, ok)
2211
+
2212
+ def _drop_log_rows(vals: dict) -> None:
2213
+ log_on = any(
2214
+ getattr(scales.get(a), "trans", None) == "log10" for a in axes
2215
+ )
2216
+ if not log_on or "frames" in vals:
2217
+ return
2218
+ ok = np.isfinite(np.asarray(vals["x"], dtype=np.float64))
2219
+ if "y" in vals:
2220
+ ok = ok & np.isfinite(np.asarray(vals["y"], dtype=np.float64))
2221
+ if "z" in vals:
2222
+ ok = ok & np.isfinite(np.asarray(vals["z"], dtype=np.float64))
2223
+ if ok.all():
2224
+ return
2225
+ if not ok.any():
2226
+ raise ValueError(
2227
+ "log scale removed every row; values must be positive"
2228
+ )
2229
+ _filter_rows(vals, ok)
2230
+
2231
+ def _vals_from_step(anim: dict) -> dict:
2232
+ parts_x = []
2233
+ parts_y = []
2234
+ spans = []
2235
+ groups = []
2236
+ cursor = 0
2237
+ for fr in anim["frames"]:
2238
+ xs = np.asarray(fr["x"], dtype=np.float64).ravel()
2239
+ ys = np.asarray(fr["y"], dtype=np.float64).ravel()
2240
+ parts_x.append(xs)
2241
+ parts_y.append(ys)
2242
+ spans.append([cursor, int(xs.size)])
2243
+ groups.append([[int(a), int(b)] for a, b in fr["groups"]])
2244
+ cursor += int(xs.size)
2245
+ all_x = np.concatenate(parts_x) if cursor else np.zeros(0)
2246
+ all_y = np.concatenate(parts_y) if cursor else np.zeros(0)
2247
+ tx = np.asarray(_absorb_position("x", "num", all_x, []), dtype=np.float64)
2248
+ ty = np.asarray(_absorb_position("y", "num", all_y, []), dtype=np.float64)
2249
+ static_i = int(anim.get("static", len(spans) - 1))
2250
+ start, count = spans[static_i]
2251
+ return {
2252
+ "x": tx[start:start + count].copy(),
2253
+ "y": ty[start:start + count].copy(),
2254
+ "frames": {
2255
+ "mode": "step",
2256
+ "nFrames": int(len(spans)),
2257
+ "spans": spans,
2258
+ "groups": groups,
2259
+ "x": tx,
2260
+ "y": ty,
2261
+ },
2262
+ }
2263
+
2264
+ def _vals_from_function_anim(anim: dict) -> dict:
2265
+ if anim.get("mode") == "step":
2266
+ return _vals_from_step(anim)
2267
+ channels = anim["channels"]
2268
+ sample = next(iter(channels.values()))
2269
+ _n, n_frames = sample.shape
2270
+ frames = {
2271
+ "mode": "tween",
2272
+ "nFrames": int(n_frames),
2273
+ "nObj": int(sample.shape[0]),
2274
+ }
2275
+ vals = {"frames": frames}
2276
+ # Sliders show every parameter at its low end. Transitions keep
2277
+ # the last frame, which is what a paused chart already showed.
2278
+ static_col = int(anim.get("static_col", -1))
2279
+ for axis_name, mat in channels.items():
2280
+ mat = np.asarray(mat, dtype=np.float64)
2281
+ flat = _absorb_position(axis_name, "num", mat.ravel(), [])
2282
+ stored = np.asarray(flat, dtype=np.float64).reshape(mat.shape)
2283
+ frames[axis_name] = stored
2284
+ vals[axis_name] = stored[:, static_col].copy()
2285
+ return vals
2286
+
2287
+ # Pass 1 — per-layer values + global scale domains.
2288
+ # Only selected columns are materialised to pandas at this boundary.
2289
+ layer_vals = []
2290
+ for geom, m in resolved:
2291
+ anim = getattr(geom, "_anim", None)
2292
+ if anim is not None:
2293
+ layer_vals.append(_vals_from_function_anim(anim))
2294
+ continue
2295
+ frame = getattr(geom, "data_override", None)
2296
+ if frame is None:
2297
+ frame = getattr(geom, "layer_data", None)
2298
+ if frame is None:
2299
+ frame = data
2300
+ if geom.kind == "box":
2301
+ # Stats frame: x + ymin/lower/middle/upper/ymax (+ optional colour).
2302
+ xcol = m["x"]
2303
+ y_stat_cols = getattr(
2304
+ geom, "_stat_y_cols", ("ymin", "lower", "middle", "upper", "ymax")
2305
+ )
2306
+ cols = [xcol, *y_stat_cols]
2307
+ frame_cols = get_columns(frame)
2308
+ if "color" in m and m["color"] in frame_cols:
2309
+ cols.append(m["color"])
2310
+ sub = materialize_columns(frame, list(dict.fromkeys(cols)))
2311
+ # dropna already applied; for box require all stat y cols present
2312
+ sub = sub.dropna(subset=[xcol, *y_stat_cols])
2313
+ vals = {}
2314
+ kind, v, cats = col_values(sub[xcol])
2315
+ box_levels = getattr(geom, "_violin_levels", None)
2316
+ if box_levels is not None and kind == "num":
2317
+ # Dodged boxes: numeric positions on the category axis.
2318
+ vals["x"] = _absorb_level_positions(v, list(box_levels), "x")
2319
+ else:
2320
+ vals["x"] = _absorb_position("x", kind, v, cats)
2321
+ for col_name in y_stat_cols:
2322
+ kind_y, v_y, cats_y = col_values(sub[col_name])
2323
+ if kind_y != "num":
2324
+ raise ValueError("geom_boxplot() y statistics must be numeric")
2325
+ vals[col_name] = _absorb_position("y", kind_y, v_y, cats_y)
2326
+ # Use middle for a generic y channel (hover / fallback).
2327
+ vals["y"] = vals["middle"]
2328
+ if "color" in m and m["color"] in sub.columns:
2329
+ kind, cv, ccats = col_values(sub[m["color"]])
2330
+ if kind == "cat" or (kind == "num" and ccats):
2331
+ if color_scale is None:
2332
+ color_scale = ["cat", _forced_levels(g, ccats)]
2333
+ else:
2334
+ color_scale[1] = list(
2335
+ dict.fromkeys(color_scale[1] + ccats)
2336
+ )
2337
+ vals["color"] = ("cat", cv, ccats)
2338
+ else:
2339
+ if color_scale is None:
2340
+ color_scale = ["num", math.inf, -math.inf]
2341
+ color_scale[1] = min(color_scale[1], float(np.nanmin(cv)))
2342
+ color_scale[2] = max(color_scale[2], float(np.nanmax(cv)))
2343
+ num_color_vals.append(np.asarray(cv, dtype=np.float64))
2344
+ vals["color"] = ("num", cv, None)
2345
+ # Outliers share the same scales.
2346
+ outliers = getattr(geom, "_outlier_frame", None)
2347
+ y_name = getattr(geom, "_y_name", "y")
2348
+ if outliers is not None and len(outliers):
2349
+ ox_kind, ox, ox_cats = col_values(outliers[xcol])
2350
+ oy_kind, oy, oy_cats = col_values(outliers[y_name])
2351
+ if box_levels is not None and ox_kind == "num":
2352
+ vals["ox"] = _absorb_level_positions(ox, list(box_levels), "x")
2353
+ else:
2354
+ vals["ox"] = _absorb_position("x", ox_kind, ox, ox_cats)
2355
+ vals["oy"] = _absorb_position("y", oy_kind, oy, oy_cats)
2356
+ if "color" in m and m["color"] in outliers.columns:
2357
+ okind, ocv, occats = col_values(outliers[m["color"]])
2358
+ if okind == "cat" or (okind == "num" and occats):
2359
+ if color_scale is None:
2360
+ color_scale = ["cat", list(occats)]
2361
+ else:
2362
+ color_scale[1] = list(
2363
+ dict.fromkeys(color_scale[1] + occats)
2364
+ )
2365
+ vals["ocolor"] = ("cat", ocv, occats)
2366
+ else:
2367
+ if color_scale is None:
2368
+ color_scale = ["num", math.inf, -math.inf]
2369
+ color_scale[1] = min(
2370
+ color_scale[1], float(np.nanmin(ocv))
2371
+ )
2372
+ color_scale[2] = max(
2373
+ color_scale[2], float(np.nanmax(ocv))
2374
+ )
2375
+ num_color_vals.append(np.asarray(ocv, dtype=np.float64))
2376
+ vals["ocolor"] = ("num", ocv, None)
2377
+ layer_vals.append(vals)
2378
+ continue
2379
+
2380
+ cols = [m[a] for a in axes if a in m] + (
2381
+ [m["color"]] if "color" in m else []
2382
+ ) + ([m["group"]] if "group" in m else [])
2383
+ if geom.kind == "point" and "size" in m:
2384
+ cols.append(m["size"])
2385
+ if geom.kind == "point" and "alpha" in m:
2386
+ cols.append(m["alpha"])
2387
+ if geom.kind == "text":
2388
+ if "label" not in m:
2389
+ raise ValueError("geom_text() requires aes(label=)")
2390
+ cols.append(m["label"])
2391
+ if geom.kind == "point" and "shape" in m:
2392
+ cols.append(m["shape"])
2393
+ if geom.kind == "line" and "linetype" in m:
2394
+ cols.append(m["linetype"])
2395
+ if transition is not None and getattr(transition, "column", None):
2396
+ cols.append(transition.column)
2397
+ sub = materialize_columns(frame, list(dict.fromkeys(cols)))
2398
+ # Inf and -Inf have no place on an axis: drop them and say so.
2399
+ position_cols = [m[a] for a in axes if a in m and m[a] in sub.columns]
2400
+ infinite = np.zeros(len(sub), dtype=bool)
2401
+ for col in dict.fromkeys(position_cols):
2402
+ if pd.api.types.is_numeric_dtype(sub[col]) and not pd.api.types.is_bool_dtype(sub[col]):
2403
+ infinite |= np.isinf(sub[col].to_numpy(dtype=np.float64, na_value=np.nan))
2404
+ if infinite.any():
2405
+ count = int(infinite.sum())
2406
+ sub = sub.loc[~infinite]
2407
+ missing_notes.append(
2408
+ f"Removed {count} {'row' if count == 1 else 'rows'} containing "
2409
+ f"non-finite values ({_layer_name(geom)})"
2410
+ )
2411
+ frame_rows_dropped = count
2412
+ else:
2413
+ frame_rows_dropped = 0
2414
+ if getattr(geom, "data_override", None) is None and frame is not None:
2415
+ # ggplot2 says when rows are dropped; plot3 used to do it silently.
2416
+ removed = n_rows(frame) - len(sub) - frame_rows_dropped
2417
+ if removed > 0:
2418
+ word = "row" if removed == 1 else "rows"
2419
+ missing_notes.append(
2420
+ f"Removed {removed} {word} containing missing values "
2421
+ f"({_layer_name(geom)})"
2422
+ )
2423
+
2424
+ if (
2425
+ transition is not None
2426
+ and geom.kind == "point"
2427
+ and not getattr(transition, "ranges", None)
2428
+ ):
2429
+ pivot = _pivot_transition(sub, m, transition)
2430
+ channels: dict[str, np.ndarray] = {}
2431
+ if "size" in m:
2432
+ skind, sv, _ = col_values(sub[m["size"]])
2433
+ if skind != "num":
2434
+ raise ValueError("aes(size=) needs a numeric column")
2435
+ sv = np.asarray(sv, dtype=np.float64)
2436
+ sv = np.where(np.isfinite(sv) & (sv >= 0), sv, np.nan)
2437
+ channels["size"] = pivot["_mat"](sv)
2438
+ if not size_label:
2439
+ size_label = str(m["size"])
2440
+ # Largest bubbles first, so later (smaller) points stay visible.
2441
+ if "size" in channels:
2442
+ with np.errstate(all="ignore"):
2443
+ score = np.nanmax(channels["size"], axis=1)
2444
+ score = np.where(np.isfinite(score), score, -np.inf)
2445
+ pivot["order"] = np.argsort(-score, kind="mergesort")
2446
+ order = pivot["order"]
2447
+ ids = [pivot["ids"][i] for i in order]
2448
+ frames = {
2449
+ "ids": ids,
2450
+ "times": pivot["times"],
2451
+ "integer": bool(pivot["integer"]),
2452
+ "kind": pivot["kind"],
2453
+ "column": pivot["column"],
2454
+ "nFrames": int(pivot["nFrames"]),
2455
+ "nObj": len(ids),
2456
+ }
2457
+ vals = {"ids": ids, "frames": frames}
2458
+ for a in axes:
2459
+ kind, raw, cats = col_values(sub[m[a]])
2460
+ mat = pivot["_mat"](raw)[order]
2461
+ flat = _absorb_position(a, kind, mat.ravel(), cats)
2462
+ frames[a] = np.asarray(flat, dtype=np.float64).reshape(mat.shape)
2463
+ # Static fallback is the last keyframe (missing stays missing).
2464
+ vals[a] = frames[a][:, -1].copy()
2465
+ if "size" in channels:
2466
+ frames["size"] = channels["size"][order]
2467
+ vals["size"] = _last_finite(frames["size"])
2468
+ if "color" in m:
2469
+ kind, cv, ccats = col_values(sub[m["color"]])
2470
+ color_mat = pivot["_mat"](cv)[order]
2471
+ frames["color"] = color_mat
2472
+ frames["color_kind"] = "cat" if (kind == "cat" or ccats) else "num"
2473
+ frames["color_cats"] = list(ccats)
2474
+ shown = _last_finite(color_mat)
2475
+ _take_color(kind, shown, ccats, into=vals)
2476
+ if frames["color_kind"] == "num":
2477
+ full = np.asarray(cv, dtype=np.float64)
2478
+ num_color_vals[-1] = full
2479
+ finite = full[np.isfinite(full)]
2480
+ if finite.size and color_scale and color_scale[0] == "num":
2481
+ color_scale[1] = min(color_scale[1], float(finite.min()))
2482
+ color_scale[2] = max(color_scale[2], float(finite.max()))
2483
+ meta = {
2484
+ "type": frames["kind"],
2485
+ "column": frames["column"],
2486
+ "nFrames": frames["nFrames"],
2487
+ "times": list(frames["times"]),
2488
+ "integer": frames["integer"],
2489
+ "ease": "smooth" if frames["kind"] == "states" else "linear",
2490
+ "duration": 12,
2491
+ }
2492
+ if transition_meta is None:
2493
+ transition_meta = meta
2494
+ elif transition_meta["times"] != meta["times"] or transition_meta["type"] != meta["type"]:
2495
+ raise ValueError(
2496
+ f"{transition_meta['type']} layers disagree on frame values"
2497
+ )
2498
+ layer_vals.append(vals)
2499
+ continue
2500
+
2501
+ vals = {}
2502
+ for a in axes:
2503
+ kind, v, cats = col_values(sub[m[a]])
2504
+ levels = (
2505
+ getattr(geom, "_violin_levels", None) if a == "x"
2506
+ else getattr(geom, "_y_levels", None) if a == "y"
2507
+ else None
2508
+ )
2509
+ if levels is not None and kind == "num":
2510
+ # Numeric positions on a categorical axis (violins, dodged
2511
+ # bars, error bars, heatmap rows, a flipped plot). Join the
2512
+ # shared category list by name, so layers agree on order.
2513
+ vals[a] = _absorb_level_positions(v, list(levels), a)
2514
+ continue
2515
+ vals[a] = _absorb_position(a, kind, v, cats)
2516
+ # Histogram: domain is full bin edges, not just bin centres.
2517
+ x_domain = getattr(geom, "_x_domain", None)
2518
+ if x_domain is not None and "x" in scales and scales["x"].kind == "num":
2519
+ dom = np.asarray([x_domain[0], x_domain[1]], dtype=np.float64)
2520
+ if getattr(scales["x"], "trans", None) == "log10":
2521
+ dom, n_bad = _log10_values(dom)
2522
+ dropped_log += n_bad
2523
+ scales["x"].widen(dom)
2524
+ # Bars / densities include the baseline at y=0 in the domain.
2525
+ # A log axis has no zero; positive bars keep the data domain.
2526
+ if (
2527
+ geom.kind in {"col", "area"}
2528
+ or getattr(geom, "_baseline_zero", False)
2529
+ ) and "y" in scales and scales["y"].kind == "num":
2530
+ if getattr(scales["y"], "trans", None) != "log10":
2531
+ scales["y"].widen(np.asarray([0.0], dtype=np.float64))
2532
+ # Violin, dodged bars, error bars: numeric x positions on a
2533
+ # categorical axis were joined to the category list above.
2534
+ if "color" in m:
2535
+ kind, cv, ccats = col_values(sub[m["color"]])
2536
+ if kind == "cat" or (kind == "num" and ccats):
2537
+ if color_scale is None:
2538
+ color_scale = ["cat", _forced_levels(g, ccats)]
2539
+ else:
2540
+ color_scale[1] = list(dict.fromkeys(color_scale[1] + ccats))
2541
+ vals["color"] = ("cat", cv, ccats)
2542
+ else:
2543
+ if color_scale is None:
2544
+ color_scale = ["num", math.inf, -math.inf]
2545
+ color_scale[1] = min(color_scale[1], float(np.nanmin(cv)))
2546
+ color_scale[2] = max(color_scale[2], float(np.nanmax(cv)))
2547
+ num_color_vals.append(np.asarray(cv, dtype=np.float64))
2548
+ vals["color"] = ("num", cv, None)
2549
+ if "group" in m:
2550
+ _, gv, gcats = col_values(sub[m["group"]])
2551
+ vals["group"] = (gv, gcats)
2552
+ if geom.kind == "point" and "shape" in m:
2553
+ vals["shape"] = _discrete_codes(sub[m["shape"]])
2554
+ if geom.kind == "line" and "linetype" in m:
2555
+ vals["linetype"] = _discrete_codes(sub[m["linetype"]])
2556
+ if geom.kind == "text":
2557
+ vals["label"] = [_label_text(v) for v in sub[m["label"]].tolist()]
2558
+ # A nudged label sits where it is drawn: the scales make room for
2559
+ # it, as ggplot2 trains them after position_nudge().
2560
+ for axis_name, shift in (("x", getattr(geom, "nudge_x", 0.0)), ("y", getattr(geom, "nudge_y", 0.0))):
2561
+ sc = scales.get(axis_name)
2562
+ if shift and sc is not None and sc.kind == "num" and axis_name in vals:
2563
+ sc.widen(np.asarray(vals[axis_name], dtype=np.float64) + float(shift))
2564
+ if geom.kind == "point" and "alpha" in m:
2565
+ akind, av, _ = col_values(sub[m["alpha"]])
2566
+ if akind != "num":
2567
+ raise ValueError("aes(alpha=) needs a numeric column")
2568
+ vals["alpha"] = np.asarray(av, dtype=np.float64)
2569
+ if not alpha_label:
2570
+ alpha_label = str(m["alpha"])
2571
+ if geom.kind == "point" and "size" in m:
2572
+ skind, sv, _ = col_values(sub[m["size"]])
2573
+ if skind != "num":
2574
+ raise ValueError("aes(size=) needs a numeric column")
2575
+ sv = np.asarray(sv, dtype=np.float64)
2576
+ # Area from zero has no place for negatives; scale_size(range=)
2577
+ # maps the whole data range, negatives included.
2578
+ ranged = getattr(getattr(g, "size_scale", None), "kind", None) == "range"
2579
+ vals["size"] = np.where(np.isfinite(sv) & (ranged | (sv >= 0)), sv, np.nan)
2580
+ if not size_label:
2581
+ size_label = str(m["size"])
2582
+ _drop_log_rows(vals)
2583
+ _limit_rows(vals, geom)
2584
+ layer_vals.append(vals)
2585
+
2586
+ # Reference lines are part of the picture: ggplot2 widens the scales so
2587
+ # geom_hline(yintercept=0) is never off the panel.
2588
+ for ref in ref_layers:
2589
+ axis = {"hline": "y", "vline": "x"}.get(ref.kind)
2590
+ if axis and axis in scales and scales[axis].kind in {"num", "dt"}:
2591
+ values = [_ref_scale_value(scales[axis], v) for v in ref.values]
2592
+ scales[axis].widen(np.asarray([v for v in values if v is not None], dtype=np.float64))
2593
+
2594
+ # expand_limits(y=0): the axes reach these values.
2595
+ for axis_name, values in (getattr(g, "expand", None) or {}).items():
2596
+ sc = scales.get(axis_name)
2597
+ if sc is not None and sc.kind in {"num", "dt"} and values:
2598
+ points = [_ref_scale_value(sc, v) for v in values]
2599
+ sc.widen(np.asarray([v for v in points if v is not None], dtype=np.float64))
2600
+ # A rug's values belong on the axes too, as ggplot2 trains its scales.
2601
+ for rug in rug_layers:
2602
+ for axis, series in _rug_series(rug, g):
2603
+ if axis in scales and scales[axis].kind in {"num", "dt"}:
2604
+ values = _rug_positions(scales[axis], series)
2605
+ scales[axis].widen(values[np.isfinite(values)])
2606
+
2607
+ for a in axes:
2608
+ scales[a].finish()
2609
+
2610
+ # Optional forced domains (facet_wrap scales="fixed").
2611
+ force = getattr(g, "_force_scales", None) or {}
2612
+ for ax, (lo, hi) in force.items():
2613
+ if ax in scales and scales[ax].kind == "num":
2614
+ lo_f, hi_f = float(lo), float(hi)
2615
+ if getattr(scales[ax], "trans", None) == "log10":
2616
+ if lo_f <= 0 or hi_f <= 0:
2617
+ raise ValueError(f"scale_{ax}_log10() limits must be positive")
2618
+ lo_f, hi_f = math.log10(lo_f), math.log10(hi_f)
2619
+ scales[ax].lo = lo_f
2620
+ scales[ax].hi = hi_f
2621
+ if scales[ax].hi <= scales[ax].lo:
2622
+ scales[ax].hi = scales[ax].lo + 1.0
2623
+
2624
+ # scale_x_continuous(limits=, breaks=, labels=), xlim(), scale_y_reverse()…
2625
+ for axis_name in ("x", "y"):
2626
+ pscale = getattr(g, f"{axis_name}scale", None)
2627
+ sc = scales.get(axis_name)
2628
+ if pscale is None or sc is None:
2629
+ continue
2630
+ if pscale.kind == "discrete" and sc.kind == "cat" and pscale.limits:
2631
+ sc.cats = list(pscale.limits)
2632
+ sc.finish()
2633
+ elif pscale.limits is not None and sc.kind in {"num", "dt"}:
2634
+ lo, hi = _limits_in_scale(pscale, sc)
2635
+ if lo is not None:
2636
+ sc.lo = lo
2637
+ if hi is not None:
2638
+ sc.hi = hi
2639
+ if sc.hi <= sc.lo:
2640
+ sc.hi = sc.lo + 1.0
2641
+ sc.custom = pscale
2642
+ if pscale.trans == "reverse" and sc.kind == "num":
2643
+ sc.lo, sc.hi = sc.hi, sc.lo
2644
+
2645
+ # A function's ylim/zlim clips the view even when samples sit inside it.
2646
+ for geom, _locked_mapping in resolved:
2647
+ lock = getattr(geom, "_axis_lock", None)
2648
+ if not lock:
2649
+ continue
2650
+ for axis_name, bounds in lock.items():
2651
+ if axis_name not in scales:
2652
+ continue
2653
+ lo, hi = float(bounds[0]), float(bounds[1])
2654
+ if hi < lo:
2655
+ lo, hi = hi, lo
2656
+ if getattr(scales[axis_name], "trans", None) == "log10":
2657
+ if lo <= 0 or hi <= 0:
2658
+ raise ValueError(
2659
+ f"scale_{axis_name}_log10() limits must be positive"
2660
+ )
2661
+ lo, hi = math.log10(lo), math.log10(hi)
2662
+ if hi <= lo:
2663
+ hi = lo + 1.0
2664
+ scales[axis_name].lo = lo
2665
+ scales[axis_name].hi = hi
2666
+
2667
+ # coord_cartesian(xlim=, ylim=): the view only; every row was used above.
2668
+ if isinstance(coord, coord_cartesian):
2669
+ for axis_name, lim in (("x", coord.xlim), ("y", coord.ylim)):
2670
+ sc = scales.get(axis_name)
2671
+ if lim is None or sc is None or sc.kind not in {"num", "dt"}:
2672
+ continue
2673
+ lo, hi = _limits_in_scale(_Limits(lim), sc)
2674
+ if lo is not None:
2675
+ sc.lo = lo
2676
+ if hi is not None:
2677
+ sc.hi = hi
2678
+ if sc.hi <= sc.lo:
2679
+ sc.hi = sc.lo + 1.0
2680
+ if getattr(sc, "custom", None) is not None and sc.custom.trans == "reverse":
2681
+ sc.lo, sc.hi = sc.hi, sc.lo
2682
+
2683
+ # Numeric colour limits: robust 2-98 percentile by default so skewed data
2684
+ # (lidar intensity) actually varies; override via scale_colour_continuous.
2685
+ num_color = None
2686
+ if color_scale is not None and color_scale[0] == "num":
2687
+ allv = np.concatenate(num_color_vals) if num_color_vals else np.array([0.0, 1.0])
2688
+ cs = g.cscale or scale_colour_continuous()
2689
+ if g.cscale is None and any(getattr(geom, "_default_ramp", None) for geom, _m in resolved):
2690
+ # A formula surface coloured by height spans its whole range.
2691
+ cs = scale_colour_continuous(limits="full")
2692
+ user_scale = getattr(g, "colour_scale", None)
2693
+ if user_scale is not None and user_scale.kind == "discrete":
2694
+ raise ValueError(
2695
+ "this colour is numeric; scale_colour_manual() and brewer are for "
2696
+ "groups. Use scale_colour_gradient(), or map a text column"
2697
+ )
2698
+ if user_scale is not None and user_scale.kind == "continuous":
2699
+ # ggplot2 maps the whole data range, unless limits are given.
2700
+ cs = scale_colour_continuous(limits=user_scale.limits or "full")
2701
+ if "color" in force:
2702
+ lo_c, hi_c = force["color"]
2703
+ lo_c, hi_c = float(lo_c), float(hi_c)
2704
+ elif cs.limits == "full":
2705
+ lo_c, hi_c = color_scale[1], color_scale[2]
2706
+ elif isinstance(cs.limits, (tuple, list)):
2707
+ lo_c, hi_c = float(cs.limits[0]), float(cs.limits[1])
2708
+ else:
2709
+ lo_c = float(np.nanpercentile(allv, 2))
2710
+ hi_c = float(np.nanpercentile(allv, 98))
2711
+ if hi_c <= lo_c:
2712
+ lo_c, hi_c = color_scale[1], color_scale[2]
2713
+ if hi_c <= lo_c:
2714
+ hi_c = lo_c + 1.0
2715
+ tf = {
2716
+ "linear": lambda a: a,
2717
+ "sqrt": lambda a: np.sqrt(np.maximum(a, 0.0)),
2718
+ "log10": lambda a: np.log10(np.maximum(a, 1e-12)),
2719
+ }[cs.trans]
2720
+ num_color = (lo_c, hi_c, cs.trans, tf)
2721
+
2722
+ # One size scale for the figure. Area: radius follows sqrt(value / max).
2723
+ size_pieces = []
2724
+ for vals in layer_vals:
2725
+ if vals.get("size") is not None:
2726
+ size_pieces.append(np.asarray(vals["size"], dtype=np.float64).ravel())
2727
+ fr = vals.get("frames")
2728
+ if fr is not None and fr.get("size") is not None:
2729
+ size_pieces.append(np.asarray(fr["size"], dtype=np.float64).ravel())
2730
+ # aes(alpha=): the data's range onto scale_alpha's opacities (0.1 to 1).
2731
+ alpha_pieces = [np.asarray(v["alpha"], dtype=np.float64) for v in layer_vals if v.get("alpha") is not None]
2732
+ alpha_map = None
2733
+ if alpha_pieces:
2734
+ all_a = np.concatenate(alpha_pieces)
2735
+ all_a = all_a[np.isfinite(all_a)]
2736
+ alpha_scale = getattr(g, "alpha_scale", None)
2737
+ a_lo, a_hi = (alpha_scale.limits if alpha_scale is not None and alpha_scale.limits
2738
+ else (float(all_a.min()), float(all_a.max())) if all_a.size else (0.0, 1.0))
2739
+ o_lo, o_hi = alpha_scale.range if alpha_scale is not None else (0.1, 1.0)
2740
+
2741
+ def alpha_map(v, a_lo=a_lo, a_hi=a_hi, o_lo=o_lo, o_hi=o_hi):
2742
+ v = np.asarray(v, dtype=np.float64)
2743
+ span = a_hi - a_lo
2744
+ t = np.clip((v - a_lo) / span, 0.0, 1.0) if span > 0 else np.ones_like(v)
2745
+ return np.where(np.isfinite(v), o_lo + t * (o_hi - o_lo), o_lo)
2746
+
2747
+ if alpha_scale is not None and alpha_scale.name:
2748
+ alpha_label = alpha_scale.name
2749
+ size_max = None
2750
+ size_units = None
2751
+ if size_pieces:
2752
+ all_s = np.concatenate(size_pieces)
2753
+ ranged = getattr(getattr(g, "size_scale", None), "kind", None) == "range"
2754
+ ok_s = all_s[np.isfinite(all_s) & (ranged | (all_s >= 0))]
2755
+ if ok_s.size == 0:
2756
+ raise ValueError("aes(size=) has no finite, non-negative values")
2757
+ size_max = float(ok_s.max())
2758
+ size_units = _bubble_max(is3d, coord)
2759
+ size_scale = getattr(g, "size_scale", None)
2760
+ size_fraction = lambda v: _area_fraction(v, size_max) # noqa: E731
2761
+ if size_scale is not None and size_scale.kind == "area" and size_scale.max_size:
2762
+ size_units = size_scale.max_size
2763
+ elif size_scale is not None and size_scale.kind == "range":
2764
+ # ggplot2's scale_size: the data's range onto (r0, r1) by area.
2765
+ lo_s, hi_s = size_scale.limits or (float(ok_s.min()), size_max)
2766
+ r0, r1 = size_scale.range
2767
+ size_units = r1
2768
+
2769
+ def size_fraction(v, lo_s=lo_s, hi_s=hi_s, r0=r0, r1=r1):
2770
+ v = np.asarray(v, dtype=np.float64)
2771
+ span = hi_s - lo_s
2772
+ t = np.where(np.isfinite(v), np.clip((v - lo_s) / span, 0.0, 1.0) if span > 0 else 1.0, np.nan)
2773
+ return np.sqrt(r0 * r0 + t * (r1 * r1 - r0 * r0)) / r1
2774
+ for vals in layer_vals:
2775
+ if vals.get("size") is not None:
2776
+ vals["size_frac"] = size_fraction(vals["size"])
2777
+ fr = vals.get("frames")
2778
+ if fr is not None and fr.get("size") is not None:
2779
+ fr["size_frac"] = size_fraction(fr["size"])
2780
+
2781
+ # Pass 2 — encode payloads per layer (quantized against the shared scales)
2782
+ # Distinct default colours so several formulas can share a legend.
2783
+ # A fill or a tangent point then copies the layer it belongs to.
2784
+ palette = theme["cat"]
2785
+ color_slot = 0
2786
+ for geom, _mapped in resolved:
2787
+ if getattr(geom, "_legend_label", None) and geom.const_color is None:
2788
+ geom.const_color = palette[color_slot % len(palette)]
2789
+ color_slot += 1
2790
+ labeled: dict[str, str] = {}
2791
+ for geom, _mapped in resolved:
2792
+ if geom.const_color is None:
2793
+ continue
2794
+ key = getattr(geom, "_color_key", None)
2795
+ label = getattr(geom, "_legend_label", None)
2796
+ if key and key not in labeled:
2797
+ labeled[key] = geom.const_color
2798
+ if label and label not in labeled:
2799
+ labeled[label] = geom.const_color
2800
+ for geom, _mapped in resolved:
2801
+ if geom.const_color is not None:
2802
+ continue
2803
+ source = getattr(geom, "_inherit_from", None)
2804
+ if source and source in labeled:
2805
+ geom.const_color = labeled[source]
2806
+
2807
+ payloads: list[tuple[str, str]] = []
2808
+ layer_specs = []
2809
+ arrows: list[dict] = []
2810
+ # One palette for the colour/fill groups: a scale_*_manual / brewer /
2811
+ # viridis_d scale, else the theme's colours extended past eight groups.
2812
+ cat_colours: list[str] = []
2813
+ if color_scale is not None and color_scale[0] == "cat":
2814
+ user_scale = getattr(g, "colour_scale", None)
2815
+ if user_scale is not None and user_scale.kind == "continuous":
2816
+ raise ValueError(
2817
+ "scale_colour_gradient() is for numbers; this colour is categorical. "
2818
+ "Use scale_colour_manual() or scale_colour_brewer()"
2819
+ )
2820
+ default_discrete = next(
2821
+ (getattr(geom, "_default_discrete") for geom, _m in resolved
2822
+ if getattr(geom, "_default_discrete", None)), None,
2823
+ )
2824
+ if user_scale is not None:
2825
+ cat_colours = user_scale.colours(list(color_scale[1]), theme["cat"])
2826
+ elif default_discrete:
2827
+ # Ordered bands (geom_density_2d_filled): viridis, as ggplot2.
2828
+ from plot3.scaling import scale_colour_viridis_d
2829
+
2830
+ cat_colours = scale_colour_viridis_d(default_discrete).colours(list(color_scale[1]), theme["cat"])
2831
+ else:
2832
+ from plot3.scaling import extend_palette
2833
+
2834
+ cat_colours = extend_palette(theme["cat"], len(color_scale[1]))
2835
+
2836
+ # Text layers are positioned through the scales like any layer (a label
2837
+ # at x="Sat" lands on Sat), then drawn as annotations by both renderers.
2838
+ text_anns: list[dict] = []
2839
+ kept_pairs = []
2840
+ for (geom, m), vals in zip(resolved, layer_vals):
2841
+ if geom.kind == "text":
2842
+ if is3d:
2843
+ raise ValueError("geom_text() is for 2D figures")
2844
+ text_anns.extend(_text_annotations(geom, vals, color_scale, theme, cat_colours))
2845
+ else:
2846
+ kept_pairs.append(((geom, m), vals))
2847
+ if len(kept_pairs) != len(resolved):
2848
+ resolved = [pair for pair, _vals in kept_pairs]
2849
+ layer_vals = [vals for _pair, vals in kept_pairs]
2850
+
2851
+ for li, ((geom, m), vals) in enumerate(zip(resolved, layer_vals)):
2852
+ n = len(vals["x"])
2853
+ order = np.arange(n)
2854
+ group_vec = None
2855
+ if "group" in vals:
2856
+ group_vec = vals["group"][0]
2857
+ elif vals.get("color") and vals["color"][0] == "cat":
2858
+ group_vec = vals["color"][1]
2859
+ if vals.get("linetype") is not None:
2860
+ # Each linetype level is its own line, as in ggplot2.
2861
+ lt_codes, lt_cats = vals["linetype"]
2862
+ group_vec = (
2863
+ lt_codes if group_vec is None
2864
+ else np.asarray(group_vec, dtype=np.float64) * (len(lt_cats) + 1) + lt_codes
2865
+ )
2866
+ # Precomputed breaks (implicit contours, clipped formulas) stay in order.
2867
+ if getattr(geom, "_groups", None) is not None:
2868
+ order = np.arange(n)
2869
+ elif geom.kind in {"line", "area"}:
2870
+ keys = []
2871
+ if group_vec is not None:
2872
+ keys.append(group_vec)
2873
+ if geom.kind == "area" or getattr(geom, "sort_x", False):
2874
+ keys.append(vals["x"])
2875
+ if keys:
2876
+ order = np.lexsort(tuple(reversed(keys)))
2877
+ # Large bubbles first so smaller ones, drawn later, stay visible.
2878
+ if (
2879
+ geom.kind == "point"
2880
+ and vals.get("size") is not None
2881
+ and "frames" not in vals
2882
+ ):
2883
+ raw_size = np.asarray(vals["size"], dtype=np.float64)
2884
+ score = np.where(
2885
+ np.isfinite(raw_size) & (raw_size >= 0), raw_size, -np.inf
2886
+ )
2887
+ order = np.argsort(-score, kind="mergesort")
2888
+ # poly keeps authoring order (closed violin contours).
2889
+
2890
+ spec_l = {
2891
+ "kind": geom.kind,
2892
+ "n": int(n),
2893
+ "alpha": geom.alpha,
2894
+ # Error bars, reference lines, and text default to the ink colour
2895
+ # (black on a light theme), as in ggplot2.
2896
+ "constColor": _hex_or_none(
2897
+ geom.const_color
2898
+ or (theme["ink"] if getattr(geom, "_ink_default", False) else None)
2899
+ ),
2900
+ }
2901
+
2902
+ if getattr(geom, "_blank", False):
2903
+ spec_l["blank"] = True # geom_blank: on the scales, not drawn
2904
+ if getattr(geom, "_polygon", False):
2905
+ # Any simple shape, concave too: renderers fill it whole rather
2906
+ # than as a violin's strip.
2907
+ spec_l["polygon"] = True
2908
+ if spec_l["constColor"] is None and not vals.get("color") and not is3d:
2909
+ # No colour of its own: black marks and grey35 bars on a light
2910
+ # theme, as ggplot2 draws them.
2911
+ spec_l["constColor"] = _hex_or_none(
2912
+ theme.get("bar") if geom.kind in {"col"} or getattr(geom, "_polygon", False)
2913
+ else theme.get("mark")
2914
+ )
2915
+ is_violin = bool(getattr(geom, "_is_violin", False))
2916
+ if (geom.kind == "box" or is_violin) and theme.get("mark") == "#000000":
2917
+ # ggplot2's boxes and violins: white inside, a dark outline.
2918
+ spec_l["plainFill"] = True
2919
+
2920
+ def _encode_channel(name: str, values: np.ndarray, axis: str):
2921
+ sc = scales[axis]
2922
+ values = np.asarray(values, dtype=np.float64)
2923
+ if not np.isfinite(values).all():
2924
+ fill = sc.lo if math.isfinite(sc.lo) else 0.0
2925
+ values = np.where(np.isfinite(values), values, fill)
2926
+ # 16-bit positions clamp to the scale. Rows past it (a
2927
+ # coord_cartesian zoom) keep their place as floats, and the
2928
+ # panel clips them.
2929
+ lo_v, hi_v = min(sc.lo, sc.hi), max(sc.lo, sc.hi)
2930
+ slack = 1e-9 * max(abs(hi_v - lo_v), 1.0)
2931
+ inside = bool(values.size == 0 or (values.min() >= lo_v - slack and values.max() <= hi_v + slack))
2932
+ enc = encode_norm(
2933
+ values,
2934
+ sc.lo,
2935
+ sc.hi,
2936
+ quantize=g.quantize and inside,
2937
+ compress=g.compress,
2938
+ )
2939
+ pid = f"p{li}{name}"
2940
+ payloads.append((pid, enc["b64"]))
2941
+ spec_l[name] = {"id": pid, "dtype": enc["dtype"]}
2942
+
2943
+ def _encode_matrix(tag: str, values: np.ndarray, lo: float, hi: float):
2944
+ flat = np.asarray(values, dtype=np.float64).ravel()
2945
+ present = np.isfinite(flat)
2946
+ fill = lo if math.isfinite(lo) else 0.0
2947
+ filled = np.where(present, flat, fill)
2948
+ enc = encode_norm(
2949
+ filled, lo, hi, quantize=g.quantize, compress=g.compress
2950
+ )
2951
+ pid = f"p{li}f{tag}"
2952
+ mid = f"p{li}f{tag}m"
2953
+ payloads.append((pid, enc["b64"]))
2954
+ payloads.append((mid, pack_u8(present.astype(np.uint8), g.compress)))
2955
+ return {"id": pid, "dtype": enc["dtype"], "mask": mid}
2956
+
2957
+ if geom.kind in {"surface", "isosurface"}:
2958
+ for a in axes:
2959
+ _encode_channel(a, vals[a][order], a)
2960
+ indices = getattr(geom, "_indices", None)
2961
+ if indices is None:
2962
+ raise ValueError(f"{geom.kind} missing triangle indices")
2963
+ flat = np.ascontiguousarray(indices.reshape(-1), dtype=np.uint32)
2964
+ pid = f"p{li}idx"
2965
+ payloads.append((pid, pack_u32(flat, g.compress)))
2966
+ spec_l["indices"] = {
2967
+ "id": pid,
2968
+ "dtype": "u32",
2969
+ "count": int(flat.size),
2970
+ }
2971
+ spec_l["wireframe"] = bool(getattr(geom, "wireframe", False))
2972
+ if getattr(geom, "_nx", None) is not None:
2973
+ spec_l["nx"] = int(geom._nx)
2974
+ spec_l["ny"] = int(geom._ny)
2975
+ if spec_l["alpha"] is None:
2976
+ spec_l["alpha"] = 0.95 if geom.kind == "surface" else 0.55
2977
+ elif geom.kind == "box":
2978
+ for name in ("x", "ymin", "lower", "middle", "upper", "ymax"):
2979
+ _encode_channel(name, vals[name][order], "x" if name == "x" else "y")
2980
+ # Keep y as middle for shared hover helpers.
2981
+ spec_l["y"] = spec_l["middle"]
2982
+ if getattr(geom, "_fill_mapped", False):
2983
+ # ggplot2: aes(fill=) fills the box; outline, whiskers, and
2984
+ # median stay dark.
2985
+ spec_l["fillMapped"] = True
2986
+ if vals.get("ox") is not None:
2987
+ n_out = len(vals["ox"])
2988
+ spec_l["nOut"] = int(n_out)
2989
+ _encode_channel("ox", vals["ox"], "x")
2990
+ _encode_channel("oy", vals["oy"], "y")
2991
+ if vals.get("ocolor"):
2992
+ ckind, cv, _ = vals["ocolor"]
2993
+ pid = f"p{li}oc"
2994
+ if ckind == "cat":
2995
+ local = vals["ocolor"][2]
2996
+ remap = {
2997
+ i: color_scale[1].index(c)
2998
+ for i, c in enumerate(local)
2999
+ }
3000
+ codes = np.array(
3001
+ [remap.get(int(c), 0) for c in cv], dtype="<u2"
3002
+ )
3003
+ payloads.append((pid, pack_u16(codes, g.compress)))
3004
+ spec_l["ocolor"] = {
3005
+ "id": pid, "dtype": "u16", "kind": "cat"
3006
+ }
3007
+ else:
3008
+ lo_c, hi_c, _trans, tf = num_color
3009
+ cvt = tf(np.clip(cv, lo_c, hi_c))
3010
+ lo_t = float(tf(np.asarray(lo_c)))
3011
+ hi_t = float(tf(np.asarray(hi_c)))
3012
+ enc = encode_norm(
3013
+ cvt, lo_t, hi_t, quantize=True, compress=g.compress
3014
+ )
3015
+ payloads.append((pid, enc["b64"]))
3016
+ spec_l["ocolor"] = {
3017
+ "id": pid, "dtype": "u16", "kind": "num"
3018
+ }
3019
+ else:
3020
+ spec_l["nOut"] = 0
3021
+ elif geom.kind not in {"surface", "isosurface"}:
3022
+ for a in axes:
3023
+ _encode_channel(a, vals[a][order], a)
3024
+
3025
+ if vals.get("color") and geom.kind not in {"box"}:
3026
+ ckind, cv, _ = vals["color"]
3027
+ pid = f"p{li}c"
3028
+ if ckind == "cat":
3029
+ # remap onto the global cat list
3030
+ local = vals["color"][2]
3031
+ remap = {i: color_scale[1].index(c) for i, c in enumerate(local)}
3032
+ codes = np.array([remap.get(int(c), 0) for c in cv[order]],
3033
+ dtype="<u2")
3034
+ payloads.append((pid, pack_u16(codes, g.compress)))
3035
+ spec_l["color"] = {"id": pid, "dtype": "u16", "kind": "cat"}
3036
+ else:
3037
+ lo_c, hi_c, _trans, tf = num_color
3038
+ cvt = tf(np.clip(cv[order], lo_c, hi_c))
3039
+ lo_t = float(tf(np.asarray(lo_c)))
3040
+ hi_t = float(tf(np.asarray(hi_c)))
3041
+ enc = encode_norm(cvt, lo_t, hi_t, quantize=True,
3042
+ compress=g.compress)
3043
+ payloads.append((pid, enc["b64"]))
3044
+ spec_l["color"] = {"id": pid, "dtype": "u16", "kind": "num"}
3045
+ elif vals.get("color") and geom.kind == "box":
3046
+ ckind, cv, _ = vals["color"]
3047
+ pid = f"p{li}c"
3048
+ if ckind == "cat":
3049
+ local = vals["color"][2]
3050
+ remap = {i: color_scale[1].index(c) for i, c in enumerate(local)}
3051
+ codes = np.array(
3052
+ [remap.get(int(c), 0) for c in cv[order]], dtype="<u2"
3053
+ )
3054
+ payloads.append((pid, pack_u16(codes, g.compress)))
3055
+ spec_l["color"] = {"id": pid, "dtype": "u16", "kind": "cat"}
3056
+ else:
3057
+ lo_c, hi_c, _trans, tf = num_color
3058
+ cvt = tf(np.clip(cv[order], lo_c, hi_c))
3059
+ lo_t = float(tf(np.asarray(lo_c)))
3060
+ hi_t = float(tf(np.asarray(hi_c)))
3061
+ enc = encode_norm(
3062
+ cvt, lo_t, hi_t, quantize=True, compress=g.compress
3063
+ )
3064
+ payloads.append((pid, enc["b64"]))
3065
+ spec_l["color"] = {"id": pid, "dtype": "u16", "kind": "num"}
3066
+
3067
+ if geom.kind in {"line", "area", "poly"}:
3068
+ preset_groups = getattr(geom, "_groups", None)
3069
+ if preset_groups is not None:
3070
+ spec_l["groups"] = [
3071
+ [int(start), int(count)] for start, count in preset_groups
3072
+ ]
3073
+ spec_l["linewidth"] = float(getattr(geom, "linewidth", 2.0))
3074
+ if geom.kind in {"area", "poly"} and spec_l["alpha"] is None:
3075
+ spec_l["alpha"] = 0.4 if geom.kind == "area" else 0.45
3076
+ elif group_vec is not None:
3077
+ gv = group_vec[order]
3078
+ cut = np.flatnonzero(np.diff(gv)) + 1
3079
+ starts = np.concatenate([[0], cut])
3080
+ counts = np.diff(np.concatenate([starts, [n]]))
3081
+ spec_l["groups"] = [
3082
+ [int(s), int(c)] for s, c in zip(starts, counts)
3083
+ ]
3084
+ else:
3085
+ spec_l["groups"] = [[0, int(n)]]
3086
+ spec_l["linewidth"] = float(getattr(geom, "linewidth", 2.0))
3087
+ if geom.kind in {"area", "poly"} and spec_l["alpha"] is None:
3088
+ spec_l["alpha"] = 0.4 if geom.kind == "area" else 0.45
3089
+ if geom.kind == "line":
3090
+ _encode_dashes(spec_l, geom, vals, order, getattr(g, "linetype_scale", None))
3091
+ if getattr(geom, "arrow", None) is not None:
3092
+ arrows.extend(_arrow_specs(geom, vals, order, spec_l, cat_colours, theme))
3093
+ if geom.kind == "area":
3094
+ scy = scales["y"]
3095
+ baseline = getattr(geom, "_baseline", None)
3096
+ if baseline is None:
3097
+ baseline = 0.0
3098
+ if scy.kind == "num":
3099
+ y_span = max(scy.hi - scy.lo, 1e-12)
3100
+ spec_l["y0"] = float(
3101
+ np.clip((float(baseline) - scy.lo) / y_span, 0.0, 1.0)
3102
+ )
3103
+ else:
3104
+ spec_l["y0"] = 0.0
3105
+ elif geom.kind in {"col", "box"}:
3106
+ # Bar/box width in normalized [0,1] x-space (ggplot2 resolution × width).
3107
+ scx = scales["x"]
3108
+ span = max(scx.hi - scx.lo, 1e-12)
3109
+ rel = float(
3110
+ getattr(
3111
+ geom,
3112
+ "width",
3113
+ 0.75 if geom.kind == "box" else 0.9,
3114
+ )
3115
+ )
3116
+ if getattr(geom, "_bar_width_data", None) is not None:
3117
+ # Absolute data width from stat (histogram binwidth).
3118
+ data_w = float(geom._bar_width_data) * rel
3119
+ elif scx.kind == "cat":
3120
+ # Discrete scale: unit spacing between categories (resolution = 1).
3121
+ data_w = 1.0 * rel
3122
+ else:
3123
+ # Continuous: data_width = resolution(x) * width (ggplot2).
3124
+ xs = np.asarray(vals["x"][order], dtype=np.float64)
3125
+ data_w = resolution(xs, zero=False) * rel
3126
+ spec_l["width"] = float(np.clip(data_w / span, 1e-4, 1.0))
3127
+ if geom.kind == "col":
3128
+ scy = scales["y"]
3129
+ if scy.kind == "num":
3130
+ y_span = max(scy.hi - scy.lo, 1e-12)
3131
+ spec_l["y0"] = float(
3132
+ np.clip((0.0 - scy.lo) / y_span, 0.0, 1.0)
3133
+ )
3134
+ else:
3135
+ spec_l["y0"] = 0.0
3136
+ else:
3137
+ spec_l["outlierSize"] = float(
3138
+ getattr(geom, "outlier_size", 3.0)
3139
+ )
3140
+ if spec_l["alpha"] is None:
3141
+ spec_l["alpha"] = 0.9
3142
+ else:
3143
+ if geom.kind == "point":
3144
+ _encode_shapes(
3145
+ spec_l, geom, vals, order, li, payloads, g.compress,
3146
+ getattr(g, "shape_scale", None),
3147
+ )
3148
+ if vals.get("alpha") is not None and alpha_map is not None:
3149
+ opacity = alpha_map(np.asarray(vals["alpha"], dtype=np.float64)[order])
3150
+ enc = encode_norm(opacity, 0.0, 1.0, quantize=g.quantize, compress=g.compress)
3151
+ pid = f"p{li}op"
3152
+ payloads.append((pid, enc["b64"]))
3153
+ spec_l["opacity"] = {"id": pid, "dtype": enc["dtype"]}
3154
+ if vals.get("size_frac") is not None:
3155
+ frac = np.asarray(vals["size_frac"], dtype=np.float64)[order]
3156
+ present = np.isfinite(frac)
3157
+ filled = np.where(present, frac, 0.0)
3158
+ enc = encode_norm(
3159
+ filled, 0.0, 1.0, quantize=g.quantize, compress=g.compress
3160
+ )
3161
+ pid = f"p{li}sz"
3162
+ payloads.append((pid, enc["b64"]))
3163
+ spec_l["size"] = {
3164
+ "id": pid,
3165
+ "dtype": enc["dtype"],
3166
+ "scale": "area",
3167
+ "max": float(size_units),
3168
+ "vmax": float(size_max),
3169
+ }
3170
+ elif getattr(geom, "size", None) is not None:
3171
+ spec_l["size"] = float(geom.size)
3172
+ elif is3d:
3173
+ mode = "scene"
3174
+ if coord is not None:
3175
+ mode = getattr(coord, "size_mode", "scene") or "scene"
3176
+ spec_l["size"] = _default_3d_point_size(n, size_mode=mode)
3177
+ zoom = float(getattr(coord, "zoom", 1.0) or 1.0)
3178
+ if mode == "scene" and zoom > 1.0:
3179
+ # A closer camera magnifies scene-sized points; keep the
3180
+ # grain the user sees at the default distance.
3181
+ spec_l["size"] = float(round(spec_l["size"] / zoom, 5))
3182
+ else:
3183
+ # pixels
3184
+ spec_l["size"] = 6.0 if n <= 2000 else (4.0 if n <= 20000 else 2.5)
3185
+ if spec_l["alpha"] is None:
3186
+ if is3d:
3187
+ # Opaque marks read sharper for lidar-style clouds (pcviz).
3188
+ spec_l["alpha"] = 1.0 if n <= 200_000 else 0.85
3189
+ else:
3190
+ spec_l["alpha"] = 0.85 if n <= 50000 else 0.6
3191
+ if spec_l["alpha"] is None:
3192
+ spec_l["alpha"] = 1.0
3193
+ frames = vals.get("frames")
3194
+ if frames and frames.get("mode") == "step":
3195
+ spec_l["frames"] = {
3196
+ "mode": "step",
3197
+ "nFrames": int(frames["nFrames"]),
3198
+ "spans": frames["spans"],
3199
+ "groups": frames["groups"],
3200
+ }
3201
+ for axis_name in ("x", "y"):
3202
+ spec_l["frames"][axis_name] = _encode_matrix(
3203
+ axis_name,
3204
+ frames[axis_name],
3205
+ scales[axis_name].lo,
3206
+ scales[axis_name].hi,
3207
+ )
3208
+ elif frames and geom.kind == "point":
3209
+ spec_l["frames"] = {
3210
+ "nFrames": int(frames["nFrames"]),
3211
+ "nObj": int(frames["nObj"]),
3212
+ }
3213
+ for a in axes:
3214
+ spec_l["frames"][a] = _encode_matrix(
3215
+ a, frames[a], scales[a].lo, scales[a].hi
3216
+ )
3217
+ if frames.get("size_frac") is not None:
3218
+ ch = _encode_matrix("sz", frames["size_frac"], 0.0, 1.0)
3219
+ ch["scale"] = "area"
3220
+ ch["max"] = float(size_units)
3221
+ ch["vmax"] = float(size_max)
3222
+ spec_l["frames"]["size"] = ch
3223
+ if frames.get("color") is not None and color_scale is not None:
3224
+ if frames["color_kind"] == "cat":
3225
+ local = frames["color_cats"]
3226
+ remap = {
3227
+ i: color_scale[1].index(c) for i, c in enumerate(local)
3228
+ }
3229
+ raw = np.asarray(frames["color"], dtype=np.float64).ravel()
3230
+ flat = np.full(raw.shape, 65535, dtype=np.uint16)
3231
+ ok = np.isfinite(raw) & (raw >= 0)
3232
+ if ok.any():
3233
+ mapped = np.array(
3234
+ [remap.get(int(c), 0) for c in raw[ok]],
3235
+ dtype=np.uint16,
3236
+ )
3237
+ flat[np.flatnonzero(ok)] = mapped
3238
+ enc = encode_codes(flat, g.compress)
3239
+ pid = f"p{li}fc"
3240
+ payloads.append((pid, enc["b64"]))
3241
+ spec_l["frames"]["color"] = {
3242
+ "id": pid, "dtype": "u16", "kind": "cat",
3243
+ }
3244
+ else:
3245
+ lo_c, hi_c, _trans, tf = num_color
3246
+ raw = np.asarray(frames["color"], dtype=np.float64)
3247
+ present = np.isfinite(raw)
3248
+ transformed = np.full(raw.shape, np.nan, dtype=np.float64)
3249
+ if present.any():
3250
+ transformed[present] = tf(np.clip(raw[present], lo_c, hi_c))
3251
+ lo_t = float(tf(np.asarray(lo_c)))
3252
+ hi_t = float(tf(np.asarray(hi_c)))
3253
+ if not math.isfinite(lo_t) or not math.isfinite(hi_t) or hi_t <= lo_t:
3254
+ lo_t, hi_t = 0.0, 1.0
3255
+ ch = _encode_matrix("c", transformed, lo_t, hi_t)
3256
+ ch["kind"] = "num"
3257
+ spec_l["frames"]["color"] = ch
3258
+ elif frames:
3259
+ spec_l["frames"] = {
3260
+ "nFrames": int(frames["nFrames"]),
3261
+ "nObj": int(frames["nObj"]),
3262
+ }
3263
+ for a in axes:
3264
+ if frames.get(a) is None:
3265
+ continue
3266
+ spec_l["frames"][a] = _encode_matrix(
3267
+ a, frames[a], scales[a].lo, scales[a].hi
3268
+ )
3269
+ if geom.kind == "point":
3270
+ labels = None
3271
+ if isinstance(vals.get("ids"), list) and len(vals["ids"]) == n:
3272
+ labels = list(vals["ids"])
3273
+ elif "group" in vals:
3274
+ gv, gcats = vals["group"]
3275
+ labels = []
3276
+ for c in np.asarray(gv, dtype=np.float64):
3277
+ if not math.isfinite(c):
3278
+ labels.append("")
3279
+ continue
3280
+ idx = int(c)
3281
+ if gcats and 0 <= idx < len(gcats):
3282
+ labels.append(str(gcats[idx]))
3283
+ elif float(c).is_integer():
3284
+ labels.append(str(idx))
3285
+ else:
3286
+ labels.append(str(c))
3287
+ if labels is not None and n <= 8000:
3288
+ spec_l["ids"] = [labels[i] for i in order]
3289
+ tip_pretty = getattr(geom, "_tip_pretty", None)
3290
+ if tip_pretty:
3291
+ tip = {"pretty": str(tip_pretty)}
3292
+ tip_latex = getattr(geom, "_tip_latex", None)
3293
+ if tip_latex:
3294
+ tip["latex"] = str(tip_latex)
3295
+ spec_l["tip"] = tip
3296
+ layer_specs.append(spec_l)
3297
+
3298
+ # color spec + legend
3299
+ cspec = {"kind": "none"}
3300
+ legend = None
3301
+ if color_scale is not None:
3302
+ if color_scale[0] == "cat":
3303
+ cats = color_scale[1]
3304
+ cspec = {"kind": "cat", "palette": cat_colours, "cats": cats}
3305
+ user_scale = getattr(g, "colour_scale", None)
3306
+ if user_scale is not None and getattr(user_scale, "identity", False):
3307
+ legend = None # the colours are the data: nothing to explain
3308
+ elif user_scale is not None:
3309
+ legend = user_scale.legend_entries(list(cats), cat_colours)
3310
+ for entry in legend:
3311
+ # The category a row stands for, whatever the row order.
3312
+ entry["ci"] = list(cats).index(entry.pop("_level"))
3313
+ else:
3314
+ legend = [{"label": c, "color": cat_colours[i], "ci": i} for i, c in enumerate(cats)]
3315
+ else:
3316
+ default_ramp = next(
3317
+ (getattr(geom, "_default_ramp") for geom, _m in resolved if getattr(geom, "_default_ramp", None)),
3318
+ # 3D clouds: viridis keeps its low end visible on any page;
3319
+ # theme_lidar brings its own green-to-violet heights.
3320
+ (theme.get("ramp3d") or "viridis") if is3d else "blue",
3321
+ )
3322
+ pal = (g.cscale.palette if g.cscale else default_ramp)
3323
+ ramp = _CONT_PALETTES.get(pal, theme["seq"])
3324
+ user_scale = getattr(g, "colour_scale", None)
3325
+ if user_scale is not None and user_scale.kind == "continuous":
3326
+ ramp = user_scale.ramp(num_color[0], num_color[1])
3327
+ cspec = {"kind": "num", "lo": num_color[0], "hi": num_color[1],
3328
+ "trans": num_color[2], "ramp": ramp}
3329
+ if legend is None:
3330
+ entries = []
3331
+ for geom, _mapped in resolved:
3332
+ label = getattr(geom, "_legend_label", None)
3333
+ if not label or not geom.const_color:
3334
+ continue
3335
+ entry = {"label": str(label), "color": geom.const_color}
3336
+ if getattr(geom, "_is_formula", False):
3337
+ entry["formula"] = True
3338
+ entry["_geom"] = geom
3339
+ segments = getattr(geom, "_legend_math", None)
3340
+ latex = getattr(geom, "_legend_latex", None)
3341
+ if segments:
3342
+ entry["math"] = segments
3343
+ elif latex:
3344
+ entry["latex"] = str(latex)
3345
+ entries.append(entry)
3346
+ if entries:
3347
+ legend = entries
3348
+
3349
+ legend, shape_legend, linetype_legend = _aux_legends(
3350
+ resolved, layer_vals, legend,
3351
+ getattr(g, "shape_scale", None), getattr(g, "linetype_scale", None),
3352
+ )
3353
+
3354
+ base_map = dict(g.mapping)
3355
+ coord_spec = _coord_spec(coord, is3d, resolved, scales)
3356
+ labs_math: dict[str, list] = {}
3357
+
3358
+ def _take(key: str, text) -> str:
3359
+ raw = "" if text is None else str(text)
3360
+ if "$" not in raw:
3361
+ return raw
3362
+ plain, segments = split_math(raw)
3363
+ if segments:
3364
+ labs_math[key] = segments
3365
+ return plain
3366
+
3367
+ formula_geoms = [
3368
+ geom for geom, _mapped in resolved if getattr(geom, "_formula_primary", False)
3369
+ ]
3370
+ raw_title = g.labs.get("title") or ""
3371
+ # One function and no title of your own: the formula is the title.
3372
+ # A shaded area, a tangent, or roots stay in the legend when they
3373
+ # have a label of their own. The curve's row would only repeat the title.
3374
+ if not str(raw_title).strip() and len(formula_geoms) == 1:
3375
+ shown = formula_geoms[0]
3376
+ # Symbolic formula, not the value list. That list stays in the legend
3377
+ # when the two differ, so the title does not become a long caption.
3378
+ title = str(
3379
+ getattr(shown, "_title_label", None)
3380
+ or getattr(shown, "_legend_label", "")
3381
+ or ""
3382
+ )
3383
+ segments = getattr(shown, "_legend_math", None)
3384
+ title_latex = getattr(shown, "_title_latex", None) or getattr(shown, "_legend_latex", None)
3385
+ if segments:
3386
+ labs_math["title"] = segments
3387
+ elif title_latex:
3388
+ labs_math["title"] = [{"text": title, "latex": str(title_latex)}]
3389
+ if legend:
3390
+ legend = [
3391
+ entry for entry in legend
3392
+ if not (entry.get("formula") and entry.get("label") == title)
3393
+ ]
3394
+ if not legend:
3395
+ legend = None
3396
+ else:
3397
+ title = _take("title", raw_title)
3398
+ legend, legend_title = _formula_legend(legend, formula_geoms, bool(str(raw_title).strip()))
3399
+ if legend:
3400
+ for entry in legend:
3401
+ entry.pop("formula", None)
3402
+ entry.pop("_geom", None)
3403
+ notes: list[str] = list(dict.fromkeys(missing_notes))
3404
+ for geom, _mapped in resolved:
3405
+ for note in getattr(geom, "_notes", None) or ():
3406
+ text = str(note)
3407
+ if text and text not in notes:
3408
+ notes.append(text)
3409
+ annotations: list[dict] = []
3410
+ for geom, _mapped in resolved:
3411
+ for ann in getattr(geom, "_annotations", None) or ():
3412
+ text = str(ann.get("text") or "")
3413
+ if not text:
3414
+ continue
3415
+ annotations.append({
3416
+ "x": float(ann["x"]),
3417
+ "y": float(ann["y"]),
3418
+ "text": text,
3419
+ "style": "area",
3420
+ })
3421
+ annotations.extend(text_anns)
3422
+ if dropped_log:
3423
+ word = "value" if dropped_log == 1 else "values"
3424
+ notes.append(f"{dropped_log} non-positive {word} omitted on a log scale")
3425
+
3426
+ alpha_legend = None
3427
+ if alpha_map is not None:
3428
+ a_lo_v, a_hi_v = alpha_map.__defaults__[0], alpha_map.__defaults__[1]
3429
+ values = [v for v in nice_ticks(a_lo_v, a_hi_v, 4) if a_lo_v <= v <= a_hi_v] or [a_lo_v, a_hi_v]
3430
+ alpha_legend = {
3431
+ "label": alpha_label or "alpha",
3432
+ "breaks": [
3433
+ {"label": label, "alpha": float(alpha_map(np.asarray([v]))[0])}
3434
+ for v, label in zip(values, fmt_ticks(values))
3435
+ ],
3436
+ }
3437
+ size_legend = None
3438
+ if size_max is not None:
3439
+ breaks = []
3440
+ size_scale = getattr(g, "size_scale", None)
3441
+ if size_scale is not None and size_scale.breaks:
3442
+ values = list(size_scale.breaks)
3443
+ elif size_scale is not None and size_scale.kind == "range":
3444
+ lo_s, hi_s = size_scale.limits or (float(ok_s.min()), size_max)
3445
+ values = [v for v in nice_ticks(lo_s, hi_s, 4) if lo_s <= v <= hi_s] or [lo_s, hi_s]
3446
+ else:
3447
+ whole = bool(np.all(np.mod(ok_s, 1.0) == 0))
3448
+ values = _size_breaks(size_max, integer=whole) or ([1.0] if whole else [])
3449
+ if size_scale is not None and size_scale.name:
3450
+ size_label = size_scale.name
3451
+ for value in values:
3452
+ frac = float(size_fraction(np.asarray([value]))[0])
3453
+ breaks.append({
3454
+ "value": float(value),
3455
+ "label": fmt_num(float(value)),
3456
+ "t": float(frac),
3457
+ })
3458
+ size_legend = {
3459
+ "label": size_label or "size",
3460
+ "vmax": float(size_max),
3461
+ "max": float(size_units),
3462
+ "breaks": breaks,
3463
+ }
3464
+
3465
+ if (
3466
+ size_legend and alpha_legend and size_legend["label"] == alpha_legend["label"]
3467
+ and [b["label"] for b in size_legend["breaks"]] == [b["label"] for b in alpha_legend["breaks"]]
3468
+ ):
3469
+ # One column for both size and alpha: one legend, as ggplot2 merges them.
3470
+ for sb, ab in zip(size_legend["breaks"], alpha_legend["breaks"]):
3471
+ sb["alpha"] = ab["alpha"]
3472
+ alpha_legend = None
3473
+ spec = {
3474
+ "v": 1,
3475
+ "is3d": is3d,
3476
+ "theme": theme,
3477
+ "labs": {
3478
+ "title": title,
3479
+ "x": _take("x", _axis_label(g, base_map, resolved, "x", is3d)),
3480
+ "y": _take("y", _axis_label(g, base_map, resolved, "y", is3d)),
3481
+ "z": _take("z", _axis_label(g, base_map, resolved, "z", is3d)) if is3d else "",
3482
+ "color": (
3483
+ _legend_title_label(
3484
+ _take(
3485
+ "color",
3486
+ g.labs.get(
3487
+ "color",
3488
+ getattr(getattr(g, "colour_scale", None), "name", None)
3489
+ or base_map.get("color") or base_map.get("fill")
3490
+ or height_title
3491
+ or next(
3492
+ (getattr(geom, "_colour_title") for geom, _m in resolved
3493
+ if getattr(geom, "_colour_title", None)),
3494
+ "",
3495
+ ),
3496
+ ),
3497
+ ),
3498
+ legend_title,
3499
+ labs_math,
3500
+ )
3501
+ if (getattr(g, "theme_options", None) or {}).get("legend_title", True)
3502
+ else ""
3503
+ ),
3504
+ "subtitle": _take("subtitle", g.labs.get("subtitle")),
3505
+ "caption": _take("caption", g.labs.get("caption")),
3506
+ "tag": _take("tag", g.labs.get("tag")),
3507
+ },
3508
+ "themeOpts": _theme_opts(g),
3509
+ "facetChild": bool(getattr(g, "_facet_child", False)) or None,
3510
+ "labsMath": labs_math or None,
3511
+ "refs": _ref_specs(ref_layers, scales, theme) or None,
3512
+ "rugs": _rug_specs(rug_layers, g, scales, theme, cspec, cat_colours) or None,
3513
+ "arrows": arrows or None,
3514
+ "shapeLegend": shape_legend,
3515
+ "linetypeLegend": linetype_legend,
3516
+ "math": bool(
3517
+ labs_math
3518
+ or any(layer.get("tip") for layer in layer_specs)
3519
+ or any(entry.get("latex") or entry.get("math") for entry in (legend or []))
3520
+ ),
3521
+ "scales": {a: scales[a].spec() for a in axes},
3522
+ "color": cspec,
3523
+ "legend": legend,
3524
+ "legendPosition": (
3525
+ "none" if getattr(g, "_facet_child", False)
3526
+ else _legend_position_spec(getattr(g, "legend_position", None))
3527
+ ),
3528
+ "sizeLegend": size_legend,
3529
+ "alphaLegend": alpha_legend,
3530
+ "transition": transition_meta,
3531
+ "slider": slider_meta,
3532
+ "layers": layer_specs,
3533
+ "gz": 1 if g.compress else 0,
3534
+ "coord": coord_spec,
3535
+ "notes": notes,
3536
+ "ann": annotations,
3537
+ }
3538
+ # geom_box3d classes beside a cloud coloured by height: the entries
3539
+ # take the class column's name and the colour bar keeps its own.
3540
+ entries_title = next(
3541
+ (geom._entries_title for geom, _m in resolved if getattr(geom, "_entries_title", None)),
3542
+ None,
3543
+ )
3544
+ if entries_title and spec.get("legend") and cspec and cspec.get("kind") == "num":
3545
+ spec["labs"]["colorBar"] = spec["labs"].get("color") or ""
3546
+ spec["labs"]["color"] = entries_title
3547
+ _apply_guides(spec, getattr(g, "guides", None) or {})
3548
+ _apply_guide_options(spec, getattr(g, "guide_options", None) or {})
3549
+ return spec, payloads
3550
+
3551
+
3552
+ def _panel_grid(n: int, ncol: int | None, nrow: int | None) -> tuple[int, int]:
3553
+ if ncol is not None and nrow is not None:
3554
+ if ncol * nrow < n:
3555
+ nrow = int(math.ceil(n / ncol))
3556
+ return int(ncol), int(nrow)
3557
+ if ncol is not None:
3558
+ return int(ncol), int(math.ceil(n / ncol))
3559
+ if nrow is not None:
3560
+ return int(math.ceil(n / nrow)), int(nrow)
3561
+ # ggplot2's wrap_dims (R's n2mfrow, turned): 3 panels in one row, 4 in
3562
+ # 2 x 2, 5 or 6 in 2 rows of 3, up to 12 in rows of 4.
3563
+ if n <= 3:
3564
+ return max(1, n), 1
3565
+ if n <= 6:
3566
+ return (n + 1) // 2, 2
3567
+ if n <= 12:
3568
+ return (n + 2) // 3, 3
3569
+ rows = int(math.ceil(math.sqrt(n)))
3570
+ return int(math.ceil(n / rows)), rows
3571
+
3572
+
3573
+ def _strip_text(facet, column, level) -> str:
3574
+ """A strip's label: the level, or the facet's labeller applied to it."""
3575
+ from plot3.geoms import strip_label
3576
+
3577
+ text = _level_text(level)
3578
+ rule = getattr(facet, "labeller", None)
3579
+ return text if rule is None else strip_label(rule, column, text)
3580
+
3581
+
3582
+ def _level_text(level) -> str:
3583
+ try:
3584
+ if level is None or bool(pd.isna(level)):
3585
+ return "NA"
3586
+ except (TypeError, ValueError):
3587
+ pass
3588
+ return str(level)
3589
+
3590
+
3591
+ def _subset(data, column, level):
3592
+
3593
+ if column is None:
3594
+ return data
3595
+ return filter_equal(data, column, None if _level_text(level) == "NA" else level)
3596
+
3597
+
3598
+ def _facet_levels(data, column) -> list:
3599
+ """Panel order: a categorical column's own order, else ggplot2's sort."""
3600
+ from plot3.table import detect_backend
3601
+
3602
+ levels = unique_levels(data, column)
3603
+
3604
+ keep = False
3605
+ if detect_backend(data) == "pandas":
3606
+ keep = isinstance(data[column].dtype, pd.CategoricalDtype)
3607
+ return ordered_levels(levels, keep_order=keep)
3608
+
3609
+
3610
+ def _facet_colour_levels(g: ggplot) -> list[str] | None:
3611
+ """Categorical colour/fill levels of the whole dataset, in scale order."""
3612
+ candidates = [g.mapping.get("color"), g.mapping.get("fill")]
3613
+ for layer in g.layers:
3614
+ mapping = getattr(layer, "mapping", None) or {}
3615
+ candidates += [mapping.get("color"), mapping.get("fill")]
3616
+ for column in candidates:
3617
+ if column and has_column(g.data, column):
3618
+ kind, _values, cats = col_values(materialize_columns(g.data, [column])[column])
3619
+ if kind == "cat":
3620
+ return list(cats)
3621
+ return None
3622
+
3623
+
3624
+ _LEGEND_GLYPH = {"circle": "●", "triangle": "▲", "square": "■", "diamond": "◆",
3625
+ "plus": "+", "cross": "×"}
3626
+
3627
+
3628
+ def _facet_key_html(entry: dict, colour: str) -> str:
3629
+ """The key in front of a facet legend label: square, symbol, or line."""
3630
+ import html as _htmlesc
3631
+
3632
+ colour = _htmlesc.escape(str(colour))
3633
+ if entry.get("shape"):
3634
+ glyph = _LEGEND_GLYPH.get(str(entry["shape"]), "●")
3635
+ return (f"<span class='sw' style='background:none;width:auto;color:{colour};"
3636
+ f"font-size:11px;line-height:9px'>{glyph}</span>")
3637
+ if entry.get("dash") is not None:
3638
+ dash = entry.get("dash") or []
3639
+ da = f" stroke-dasharray='{' '.join(str(d * 1.6) for d in dash)}'" if dash else ""
3640
+ return (f"<svg width='14' height='9' style='margin-right:4px;vertical-align:middle'>"
3641
+ f"<line x1='0' y1='4.5' x2='14' y2='4.5' stroke='{colour}' stroke-width='1.6'{da}/></svg>")
3642
+ return f"<span class='sw' style='background:{colour}'></span>"
3643
+
3644
+
3645
+ def _facet_legend_html(spec: dict, theme: dict) -> str:
3646
+ """The legends of a faceted HTML figure, drawn once beside the panels:
3647
+ colour (keys with shapes or dashes), a colour bar, shape, line type,
3648
+ and size, as one panel draws them."""
3649
+ import html as _htmlesc
3650
+
3651
+ esc = _htmlesc.escape
3652
+ ink = theme["ink"]
3653
+ blocks = []
3654
+ labs = spec.get("labs") or {}
3655
+ title = str(labs.get("color") or "")
3656
+ head = f"<b style='color:{ink}'>{esc(title)}</b>" if title else ""
3657
+ entries = spec.get("legend") or []
3658
+ color = spec.get("color") or {}
3659
+ if entries:
3660
+ rows = "".join(
3661
+ f"<div>{_facet_key_html(e, e.get('color'))}{esc(str(e.get('label')))}</div>"
3662
+ for e in entries
3663
+ )
3664
+ blocks.append(head + rows)
3665
+ elif color.get("kind") == "num" and color.get("guide") is not False and color.get("ramp"):
3666
+ stops = ",".join(esc(str(c)) for c in color["ramp"])
3667
+ lo, hi = float(color.get("lo", 0.0)), float(color.get("hi", 1.0))
3668
+ blocks.append(
3669
+ f"{head}<div style='height:8px;width:110px;border-radius:4px;margin-top:3px;"
3670
+ f"background:linear-gradient(90deg,{stops})'></div>"
3671
+ f"<div style='display:flex;justify-content:space-between'><span>{lo:.3g}</span>"
3672
+ f"<span>{hi:.3g}</span></div>"
3673
+ )
3674
+ for key in ("shapeLegend", "linetypeLegend"):
3675
+ legend = spec.get(key)
3676
+ if not legend:
3677
+ continue
3678
+ rows = "".join(
3679
+ f"<div>{_facet_key_html(e, ink)}{esc(str(e.get('label')))}</div>"
3680
+ for e in legend.get("entries") or []
3681
+ )
3682
+ blocks.append(f"<b style='color:{ink}'>{esc(str(legend.get('label') or ''))}</b>{rows}")
3683
+ size = spec.get("sizeLegend")
3684
+ if size and size.get("breaks"):
3685
+ rows = "".join(
3686
+ f"<div><span style='display:inline-block;vertical-align:middle;margin-right:6px;"
3687
+ f"border-radius:50%;background:{theme['ink2']};width:{max(6, round(b.get('t', 0) * 26))}px;"
3688
+ f"height:{max(6, round(b.get('t', 0) * 26))}px'></span>{esc(str(b.get('label')))}</div>"
3689
+ for b in size["breaks"]
3690
+ )
3691
+ blocks.append(f"<b style='color:{ink}'>{esc(str(size.get('label') or 'size'))}</b>{rows}")
3692
+ alpha = spec.get("alphaLegend")
3693
+ if alpha and alpha.get("breaks"):
3694
+ rows = "".join(
3695
+ f"<div><span style='display:inline-block;vertical-align:middle;margin-right:6px;"
3696
+ f"border-radius:50%;width:10px;height:10px;background:{ink};opacity:{b.get('alpha', 1)}'>"
3697
+ f"</span>{esc(str(b.get('label')))}</div>"
3698
+ for b in alpha["breaks"]
3699
+ )
3700
+ blocks.append(f"<b style='color:{ink}'>{esc(str(alpha.get('label') or 'alpha'))}</b>{rows}")
3701
+ if not blocks:
3702
+ return ""
3703
+ return "<div id='flegend'>" + "<div style='height:6px'></div>".join(blocks) + "</div>"
3704
+
3705
+
3706
+ def facet_cells(g: ggplot) -> dict:
3707
+ """Panels of a faceted figure, for the HTML viewer and for ggsave.
3708
+
3709
+ ``cells`` lists ``{"row", "col", "fig", "strip"}``; ``fig`` is None for
3710
+ an empty facet_grid combination. facet_wrap puts each label on its own
3711
+ panel (``strip``); facet_grid uses ``col_strips`` above the top row and
3712
+ ``row_strips`` to the right, as in ggplot2.
3713
+ """
3714
+ from plot3.geoms import facet_grid as _facet_grid
3715
+
3716
+ facet = g.facet
3717
+ if g.data is None:
3718
+ raise ValueError("ggplot has no data")
3719
+ force = _global_numeric_domains(g) if facet.scales == "fixed" else {}
3720
+ colour_levels = _facet_colour_levels(g)
3721
+
3722
+ def child(panel):
3723
+ # Panels draw only data; the figure draws title, legend, axis titles.
3724
+ panel._facet_child = True
3725
+ if colour_levels:
3726
+ panel._force_color_levels = colour_levels
3727
+ if force:
3728
+ panel._force_scales = force
3729
+ return panel
3730
+ header = {
3731
+ key: g.labs.get(key) for key in ("title", "subtitle", "caption", "tag") if g.labs.get(key)
3732
+ }
3733
+ cells: list[dict] = []
3734
+ if not isinstance(facet, _facet_grid):
3735
+ column = facet.variable
3736
+ if not has_column(g.data, column):
3737
+ raise ColumnNotFound([column], g.data)
3738
+ levels = _facet_levels(g.data, column)
3739
+ if not levels:
3740
+ raise ValueError("facet_wrap() found no panel levels")
3741
+ ncol, nrow = _panel_grid(len(levels), facet.ncol, facet.nrow)
3742
+ for index, level in enumerate(levels):
3743
+ label = _strip_text(facet, column, level)
3744
+ panel = child(_clone_ggplot_with_data(g, _subset(g.data, column, level)))
3745
+ panel.labs = {
3746
+ k: v for k, v in panel.labs.items()
3747
+ if k not in {"title", "subtitle", "caption", "tag"}
3748
+ }
3749
+ row, col = divmod(index, ncol)
3750
+ cells.append({"row": row, "col": col, "fig": panel, "strip": label})
3751
+ return {"ncol": ncol, "nrow": nrow, "cells": cells, "col_strips": None,
3752
+ "row_strips": None, "header": header, "kind": "wrap"}
3753
+
3754
+ for column in (facet.rows, facet.cols):
3755
+ if column is not None and not has_column(g.data, column):
3756
+ raise ColumnNotFound([column], g.data)
3757
+ row_levels = _facet_levels(g.data, facet.rows) if facet.rows else [None]
3758
+ col_levels = _facet_levels(g.data, facet.cols) if facet.cols else [None]
3759
+ nrow, ncol = len(row_levels), len(col_levels)
3760
+ for r, row_level in enumerate(row_levels):
3761
+ rows_data = _subset(g.data, facet.rows, row_level)
3762
+ for c, col_level in enumerate(col_levels):
3763
+ piece = _subset(rows_data, facet.cols, col_level)
3764
+ if n_rows(piece) == 0:
3765
+ cells.append({"row": r, "col": c, "fig": None, "strip": None})
3766
+ continue
3767
+ panel = child(_clone_ggplot_with_data(g, piece))
3768
+ panel.labs = {
3769
+ k: v for k, v in panel.labs.items()
3770
+ if k not in {"title", "subtitle", "caption", "tag"}
3771
+ }
3772
+ cells.append({"row": r, "col": c, "fig": panel, "strip": None})
3773
+ return {
3774
+ "ncol": ncol,
3775
+ "nrow": nrow,
3776
+ "cells": cells,
3777
+ "col_strips": [_strip_text(facet, facet.cols, v) for v in col_levels] if facet.cols else None,
3778
+ "row_strips": [_strip_text(facet, facet.rows, v) for v in row_levels] if facet.rows else None,
3779
+ "header": header,
3780
+ "kind": "grid",
3781
+ }
3782
+
3783
+
3784
+ def _clone_ggplot_with_data(g: ggplot, data) -> ggplot:
3785
+ import copy
3786
+
3787
+ from plot3.table import detect_backend
3788
+
3789
+ out = copy.copy(g)
3790
+ out.data = data
3791
+ out.backend = detect_backend(data) if data is not None else None
3792
+ out.layers = list(g.layers)
3793
+ out.labs = dict(g.labs)
3794
+ out.facet = None # panels are leaf plots
3795
+ out.mapping = g.mapping
3796
+ out.cscale = g.cscale
3797
+ out.stat_density_3d = getattr(g, "stat_density_3d", None)
3798
+ out.coord = getattr(g, "coord", None)
3799
+ out._force_scales = None
3800
+ # keep theme/height/quantize
3801
+ return out
3802
+
3803
+
3804
+ def build_doc(g: ggplot) -> str:
3805
+ """Build a standalone HTML document for *g*.
3806
+
3807
+ Single-panel figures go through :func:`plot3.payload.build_payload` then
3808
+ :func:`plot3.payload.render_payload`. Faceted figures assemble a grid of
3809
+ per-panel documents (each panel uses the same payload path).
3810
+ """
3811
+ facet = getattr(g, "facet", None)
3812
+ if facet is not None:
3813
+ return _build_doc_faceted(g, facet)
3814
+
3815
+ from plot3.payload import build_payload, render_payload
3816
+
3817
+ return render_payload(build_payload(g), log=True)
3818
+
3819
+
3820
+ def _global_numeric_domains(g: ggplot) -> dict[str, tuple[float, float]]:
3821
+ """Axis/colour domains from the full faceted dataset (for scales='fixed')."""
3822
+ if g.data is None:
3823
+ return {}
3824
+ mapping = dict(g.mapping)
3825
+ cols: dict[str, str] = {}
3826
+ for ax in ("x", "y", "z"):
3827
+ if ax in mapping:
3828
+ cols[ax] = mapping[ax]
3829
+ if "color" in mapping:
3830
+ cols["color"] = mapping["color"]
3831
+ domains: dict[str, tuple[float, float]] = {}
3832
+ data_cols = get_columns(g.data)
3833
+ for ax, col in cols.items():
3834
+ if col not in data_cols:
3835
+ continue
3836
+ # Materialise only this column at the domain boundary.
3837
+ sub = materialize_columns(g.data, [col])
3838
+ series = pd.to_numeric(sub[col], errors="coerce").dropna()
3839
+ if series.empty:
3840
+ continue
3841
+ lo, hi = float(series.min()), float(series.max())
3842
+ if hi <= lo:
3843
+ hi = lo + 1.0
3844
+ domains[ax] = (lo, hi)
3845
+ return domains
3846
+
3847
+
3848
+ def _build_doc_faceted(g: ggplot, facet) -> str:
3849
+ """facet_wrap / facet_grid as a CSS grid of independent panel documents."""
3850
+ import html as _htmlesc
3851
+
3852
+ layout = facet_cells(g)
3853
+ ncol, nrow = layout["ncol"], layout["nrow"]
3854
+ theme = _figure_theme(g)
3855
+ col_strips = layout.get("col_strips")
3856
+ row_strips = layout.get("row_strips")
3857
+ by_pos = {(c["row"], c["col"]): c for c in layout["cells"]}
3858
+ first = next(c["fig"] for c in layout["cells"] if c["fig"] is not None)
3859
+ first_spec, _pairs = build_spec(first)
3860
+ shared_x = str((first_spec.get("labs") or {}).get("x") or "")
3861
+ shared_y = str((first_spec.get("labs") or {}).get("y") or "")
3862
+ legend_html = _facet_legend_html(first_spec, theme)
3863
+ cells: list[str] = []
3864
+ total_kb = 0
3865
+ count = 0
3866
+ esc = _htmlesc.escape
3867
+ # facet_grid: a strip row above the panels and a strip column to their right.
3868
+ if col_strips:
3869
+ cells += [f"<div class='cstrip'>{esc(t)}</div>" for t in col_strips]
3870
+ if row_strips:
3871
+ cells.append("<div></div>")
3872
+ for row in range(nrow):
3873
+ for col in range(ncol):
3874
+ cell = by_pos.get((row, col))
3875
+ if cell is None or cell["fig"] is None:
3876
+ cells.append("<div class='empty'></div>")
3877
+ continue
3878
+ label = cell.get("strip")
3879
+ try:
3880
+ panel_html = build_doc(cell["fig"])
3881
+ except Exception as exc:
3882
+ panel_html = (
3883
+ "<!doctype html><html><body style='font:12px system-ui;"
3884
+ f"color:#888;padding:12px'>panel {esc(str(label or ''))}: "
3885
+ f"{esc(str(exc))}</body></html>"
3886
+ )
3887
+ total_kb += len(panel_html) // 1024
3888
+ count += 1
3889
+ strip = f"<div class='plab'>{esc(label)}</div>" if label else ""
3890
+ cells.append(
3891
+ "<div class='panel'>" + strip
3892
+ + f"<iframe srcdoc=\"{esc(panel_html, quote=True)}\" title=\"panel\"></iframe></div>"
3893
+ )
3894
+ if row_strips:
3895
+ cells.append(f"<div class='rstrip'><span>{esc(row_strips[row])}</span></div>")
3896
+
3897
+ header = layout.get("header") or {}
3898
+ def text(key: str) -> str:
3899
+ raw = str(header.get(key, "") or "")
3900
+ if "$" in raw:
3901
+ raw, _segments = split_math(raw)
3902
+ return esc(raw)
3903
+ columns = f"repeat({ncol},minmax(0,1fr))" + (" 26px" if row_strips else "")
3904
+ rows = ("24px " if col_strips else "") + f"repeat({nrow},minmax(0,1fr))"
3905
+ tag = f"<b style='margin-right:8px'>{text('tag')}</b>" if header.get("tag") else ""
3906
+ subtitle = f"<div id='fsub'>{text('subtitle')}</div>" if header.get("subtitle") else ""
3907
+ caption = f"<div id='fcap'>{text('caption')}</div>" if header.get("caption") else ""
3908
+ doc = f"""<!doctype html>
3909
+ <html><head><meta charset="utf-8"><style>
3910
+ html,body{{margin:0;height:100%;background:{theme["surface"]};color:{theme["ink"]};
3911
+ font:12px system-ui,-apple-system,"Segoe UI",sans-serif}}
3912
+ #wrap{{box-sizing:border-box;height:100%;padding:8px;display:flex;flex-direction:column}}
3913
+ #ftitle{{font-size:14px;font-weight:600;margin:0 4px 2px}}
3914
+ #fsub{{font-size:12px;color:{theme["ink2"]};margin:0 4px 6px}}
3915
+ #fcap{{font-size:11px;color:{theme["muted"]};text-align:right;margin:4px 4px 0}}
3916
+ #body{{flex:1;min-height:0;display:flex;gap:6px}}
3917
+ #main{{flex:1;min-width:0;display:flex;flex-direction:column}}
3918
+ #ytitle{{width:16px;display:flex;align-items:center;justify-content:center;color:{theme["ink2"]}}}
3919
+ #ytitle span{{writing-mode:vertical-rl;transform:rotate(180deg)}}
3920
+ #xtitle{{text-align:center;color:{theme["ink2"]};padding-top:4px}}
3921
+ #flegend{{align-self:flex-start;margin-top:24px;padding:6px 9px;border:1px solid {theme["grid"]};
3922
+ border-radius:6px;font-size:11px;line-height:1.7;color:{theme["ink2"]}}}
3923
+ #flegend .sw{{display:inline-block;width:9px;height:9px;border-radius:2px;margin-right:6px}}
3924
+ #grid{{flex:1;min-height:0;display:grid;gap:8px;
3925
+ grid-template-columns:{columns};grid-template-rows:{rows}}}
3926
+ .panel{{min-height:0;min-width:0;display:flex;flex-direction:column;
3927
+ border:1px solid {theme["axis"]};border-radius:6px;overflow:hidden;
3928
+ background:{theme["surface"]}}}
3929
+ .plab{{padding:4px 8px;font-size:11px;color:{theme["ink2"]};
3930
+ border-bottom:1px solid {theme["grid"]}}}
3931
+ .cstrip,.rstrip{{background:{theme["grid"]};color:{theme["ink2"]};font-size:11px;
3932
+ font-weight:600;display:flex;align-items:center;justify-content:center;border-radius:4px}}
3933
+ .rstrip span{{writing-mode:vertical-rl}}
3934
+ .panel iframe{{flex:1;width:100%;border:0;background:{theme["surface"]}}}
3935
+ </style></head><body><div id="wrap">
3936
+ <div id="ftitle">{tag}{text("title")}</div>{subtitle}
3937
+ <div id="body"><div id="ytitle"><span>{esc(shared_y)}</span></div>
3938
+ <div id="main"><div id="grid">{"".join(cells)}</div><div id="xtitle">{esc(shared_x)}</div></div>
3939
+ {legend_html}</div>{caption}
3940
+ </div></body></html>"""
3941
+ # Sizes are for whoever embeds the HTML (a slide's ~1.8 MB cap, say):
3942
+ # PLOT3_VERBOSE=1 prints them, and a notebook cell stays just the plot.
3943
+ if os.environ.get("PLOT3_VERBOSE", "").strip() == "1":
3944
+ kind = "facet_grid" if layout.get("kind") == "grid" else "facet_wrap"
3945
+ print(f"plot3: {kind} {count} panel(s) in {nrow}x{ncol} ~{total_kb:,} KB portable HTML")
3946
+ if total_kb > 1500:
3947
+ print("plot3: the faceted figure is over 1.5 MB of HTML")
3948
+ return doc