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/geoms.py ADDED
@@ -0,0 +1,2558 @@
1
+ """Grammar objects: aes, geoms, labs, colour scales, theme helpers."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ import numpy as np
8
+ import pandas as pd
9
+
10
+ from plot3.themes import _CONT_PALETTES, _THEMES
11
+
12
+
13
+ def _as_column_name(value):
14
+ """Coerce aesthetic values to column-name strings when possible.
15
+
16
+ Jupyter R-style masking already turns bare names / backticks into strings.
17
+ This is a small runtime safety net for escaped sentinels and objects that
18
+ expose a column name (e.g. tidy3 ``col("x")``).
19
+
20
+ Integer aesthetics (``0``, ``1``, …) are kept as decimal strings so NumPy
21
+ array columns can be addressed with ``aes(x=0, y=1, z=2)``.
22
+ """
23
+ if value is None or isinstance(value, str):
24
+ return value
25
+ # Positional columns for ArrayTable / integer-named frames (not bool).
26
+ if isinstance(value, (int, np.integer)) and not isinstance(value, (bool, np.bool_)):
27
+ return str(int(value))
28
+ name = getattr(value, "name", None)
29
+ if isinstance(name, str) and name:
30
+ return name
31
+ # tidy3 / polars expr sometimes use meta_output_name or similar
32
+ for attr in ("meta_output_name", "column", "col_name"):
33
+ fn = getattr(value, attr, None)
34
+ if callable(fn):
35
+ try:
36
+ out = fn()
37
+ if isinstance(out, str) and out:
38
+ return out
39
+ except Exception:
40
+ pass
41
+ elif isinstance(fn, str) and fn:
42
+ return fn
43
+ return value
44
+
45
+
46
+ class aes(dict):
47
+ """Aesthetic mapping: aes(x=, y=, z=, colour=/color=, fill=, size=, group=,
48
+ ymin=, ymax=).
49
+
50
+ ``fill`` colours the inside of filled shapes (bars, boxes, violins,
51
+ ribbons, densities, surfaces), as in ggplot2. Points, lines, and error
52
+ bars use ``colour`` and ignore ``fill``, so ``geom_col(aes(fill=g))``
53
+ with ``geom_errorbar()`` gives coloured bars and black error bars.
54
+
55
+ ``size`` maps a numeric column onto point area (radius follows the square
56
+ root), so a bubble chart reads population as area. A constant
57
+ ``geom_point(size=)`` is still one size for the whole layer.
58
+
59
+ ``xmin``/``xmax``/``xend``/``yend`` place rectangles and segments, and
60
+ ``sample`` is the column a Q-Q plot compares with a distribution.
61
+ ``ymin`` and ``ymax`` are the ends of an error bar or ribbon. ``label``
62
+ is the text of ``geom_text``. ``shape`` picks point symbols and
63
+ ``linetype`` dash patterns per group.
64
+
65
+ ``group`` is the identity of an object. Lines use it to split series.
66
+ ``transition_time`` uses it to match the same object across frames
67
+ (a country, for example).
68
+
69
+ In Jupyter / SolveIt with R-style masking (default), bare names and
70
+ backticks work like ggplot2::
71
+
72
+ aes(x=wt, y=mpg, colour=cyl)
73
+ aes(x=`First Name`, y=`Age (%)`)
74
+
75
+ In plain ``.py`` files, use strings: ``aes(x="wt", y="mpg")``.
76
+
77
+ For 2D NumPy arrays, use integer positions (stored as ``\"0\"``, ``\"1\"``, …)::
78
+
79
+ ggplot(points, aes(x=0, y=1, z=2, colour=3)) + geom_point3d()
80
+ """
81
+
82
+ def __init__(
83
+ self,
84
+ x=None,
85
+ y=None,
86
+ z=None,
87
+ color=None,
88
+ colour=None,
89
+ fill=None,
90
+ size=None,
91
+ group=None,
92
+ ymin=None,
93
+ ymax=None,
94
+ label=None,
95
+ shape=None,
96
+ linetype=None,
97
+ xmin=None,
98
+ xmax=None,
99
+ xend=None,
100
+ yend=None,
101
+ sample=None,
102
+ width=None,
103
+ height=None,
104
+ length=None,
105
+ angle=None,
106
+ alpha=None,
107
+ ):
108
+ super().__init__()
109
+ colour_value = color if color is not None else colour
110
+ for k, v in (("x", x), ("y", y), ("z", z),
111
+ ("color", colour_value),
112
+ ("fill", fill),
113
+ ("size", size),
114
+ ("group", group),
115
+ ("ymin", ymin),
116
+ ("ymax", ymax),
117
+ ("label", label),
118
+ ("shape", shape),
119
+ ("linetype", linetype),
120
+ ("xmin", xmin),
121
+ ("xmax", xmax),
122
+ ("xend", xend),
123
+ ("yend", yend),
124
+ ("sample", sample),
125
+ ("width", width),
126
+ ("height", height),
127
+ ("length", length),
128
+ ("angle", angle),
129
+ ("alpha", alpha)):
130
+ if v is not None:
131
+ self[k] = _as_column_name(v)
132
+
133
+
134
+ def _warn_unknown(geom, params: dict) -> None:
135
+ """ggplot2's "Ignoring unknown parameters", with a guess at the intended one."""
136
+ import difflib
137
+ import inspect
138
+ import warnings
139
+
140
+ known = {"mapping", "color", "colour", "alpha", "data"}
141
+ for klass in type(geom).__mro__:
142
+ init = klass.__dict__.get("__init__")
143
+ if init is not None:
144
+ known.update(
145
+ name for name, par in inspect.signature(init).parameters.items()
146
+ if par.kind in (par.KEYWORD_ONLY, par.POSITIONAL_OR_KEYWORD)
147
+ )
148
+ name = type(geom).__name__
149
+ for key in params:
150
+ guess = difflib.get_close_matches(key, sorted(known - {"self", "mapping"}), n=1)
151
+ hint = f" (did you mean {guess[0]}?)" if guess else ""
152
+ warnings.warn(f"Ignoring unknown parameter in {name}(): {key}{hint}", stacklevel=4)
153
+
154
+
155
+ class _Geom:
156
+ kind = ""
157
+ sort_x = False
158
+ # geom_function takes formula coefficients (a=2) as keywords.
159
+ _takes_any_keyword = False
160
+
161
+ def __init__(self, mapping: aes | None = None, *, color=None, colour=None,
162
+ alpha=None, data=None, **params):
163
+ self.mapping = mapping or aes()
164
+ self.const_color = color if color is not None else colour
165
+ self.alpha = alpha
166
+ # A layer's own rows (ggplot2's geom_rect(data = periods, ...)).
167
+ self.layer_data = data
168
+ self.params = params
169
+ if params and not self._takes_any_keyword and type(self).__name__ != "_Geom":
170
+ _warn_unknown(self, params)
171
+
172
+
173
+ class geom_point(_Geom):
174
+ """Scatter points.
175
+
176
+ In **2D**, a constant ``size`` is pixels. In **3D** (when ``aes(z=...)``
177
+ is set), a constant ``size`` is scene units with distance attenuation
178
+ (unit-cube space after encoding). Prefer :class:`geom_point3d` for
179
+ explicit 3D intent. ``aes(size=)`` maps a column to area instead.
180
+
181
+ When ``size`` is omitted in 3D, a density-aware default is chosen
182
+ (pcviz-like fine points on dense clouds).
183
+ """
184
+
185
+ kind = "point"
186
+
187
+ def __init__(self, mapping=None, *, size=None, shape=None, position="identity", **kw):
188
+ super().__init__(mapping, **kw)
189
+ self.size = size
190
+ self.shape = None if shape is None else shape_name(shape)
191
+ # "jitter", position_jitter(), position_jitterdodge(), position_nudge().
192
+ self.position = position
193
+
194
+
195
+ class geom_point3d(geom_point):
196
+ """3D scatter / point-cloud marks (same ``kind`` as :class:`geom_point`).
197
+
198
+ Use with ``aes(x=, y=, z=)``. Under default :class:`coord_3d`
199
+ (``size_mode="scene"``), size is in **unit-cube scene units** with
200
+ distance attenuation — comparable to pcviz's metric sizing after
201
+ plot3 normalizes axes into a unit cube.
202
+
203
+ Parameters
204
+ ----------
205
+ size:
206
+ Point diameter in scene units. Default ``None`` picks a small,
207
+ density-aware size (roughly pcviz ``size=0.06`` m on a ~50 m cloud).
208
+ Override for artistic control, e.g. ``size=0.002``.
209
+ """
210
+
211
+ def __init__(self, mapping=None, *, size=None, **kw):
212
+ # None → build_spec density-aware default (not a hard-coded 0.01 blob).
213
+ super().__init__(mapping, size=size, **kw)
214
+
215
+
216
+ class geom_box3d(_Geom):
217
+ """Wireframe 3D boxes, one per row: detections in a lidar scene.
218
+
219
+ ``aes(x=, y=, z=)`` is the box centre, ``length`` its size along its
220
+ heading, ``width`` across it, ``height`` up z, and ``angle`` the heading
221
+ (yaw) in radians about z, as nuScenes and KITTI store boxes.
222
+ ``colour`` names a class column: each class gets its own colour and a
223
+ legend entry, separate from a point cloud coloured by height.
224
+ """
225
+
226
+ kind = "box3d"
227
+
228
+ def __init__(self, mapping=None, *, linewidth=1.5, **kw):
229
+ super().__init__(mapping, **kw)
230
+ self.linewidth = float(linewidth)
231
+
232
+
233
+ class coord_3d:
234
+ """3D coordinate system options for orbit-view figures.
235
+
236
+ Parameters
237
+ ----------
238
+ aspect:
239
+ ``"auto"`` (default) keeps relative axis spans, except that a tall z
240
+ is shortened to twice the wider horizontal side. ``"data"`` keeps
241
+ true proportions always (lidar). ``"equal"`` forces a unit cube.
242
+ size_mode:
243
+ ``"scene"`` — point size attenuates with distance (lidar).
244
+ ``"screen"`` — constant pixel size.
245
+ max_points:
246
+ If set, deterministically stride-subsample rows when building so huge
247
+ clouds stay interactive in HTML.
248
+ elev, azim:
249
+ Where the camera starts, in degrees, as matplotlib's ``view_init``:
250
+ ``elev`` above the x-y plane, ``azim`` around z from the +x axis.
251
+ The default is about ``elev=26, azim=-57``. ``azim=180, elev=25``
252
+ looks along +x from behind, the chase view of a driving scene.
253
+ zoom:
254
+ Above 1 moves the camera closer (2 is twice as close).
255
+ """
256
+
257
+ def __init__(
258
+ self,
259
+ *,
260
+ aspect: str = "auto",
261
+ size_mode: str = "scene",
262
+ max_points: int | None = None,
263
+ elev: float | None = None,
264
+ azim: float | None = None,
265
+ zoom: float = 1.0,
266
+ ):
267
+ if aspect not in {"auto", "data", "equal"}:
268
+ raise ValueError("aspect must be 'auto', 'data', or 'equal'")
269
+ if size_mode not in {"scene", "screen"}:
270
+ raise ValueError("size_mode must be 'scene' or 'screen'")
271
+ if max_points is not None and int(max_points) < 1:
272
+ raise ValueError("max_points must be positive")
273
+ self.aspect = aspect
274
+ self.size_mode = size_mode
275
+ self.max_points = None if max_points is None else int(max_points)
276
+ if elev is not None and not -90.0 <= float(elev) <= 90.0:
277
+ raise ValueError("coord_3d(elev=) is an angle from -90 to 90 degrees")
278
+ if not float(zoom) > 0:
279
+ raise ValueError("coord_3d(zoom=) must be positive")
280
+ self.elev = None if elev is None else float(elev)
281
+ self.azim = None if azim is None else float(azim)
282
+ self.zoom = float(zoom)
283
+
284
+ def to_spec(self) -> dict:
285
+ spec = {
286
+ "aspect": self.aspect,
287
+ "sizeMode": self.size_mode,
288
+ "maxPoints": self.max_points,
289
+ }
290
+ if self.elev is not None or self.azim is not None or self.zoom != 1.0:
291
+ spec["camera"] = camera_spec(self.elev, self.azim, self.zoom)
292
+ return spec
293
+
294
+
295
+ # The viewer's opening direction, as elevation and azimuth in degrees.
296
+ _DEFAULT_ELEV = math.degrees(math.asin(0.5 / math.sqrt(0.55**2 + 0.85**2 + 0.5**2)))
297
+ _DEFAULT_AZIM = math.degrees(math.atan2(-0.85, 0.55))
298
+
299
+
300
+ def camera_spec(elev=None, azim=None, zoom=1.0) -> dict:
301
+ """Unit vector from the box centre toward the camera, and the zoom."""
302
+ # Straight down leaves no "up" for the orbit; stop just short of it.
303
+ e = math.radians(max(-89.5, min(89.5, _DEFAULT_ELEV if elev is None else float(elev))))
304
+ a = math.radians(_DEFAULT_AZIM if azim is None else float(azim))
305
+ return {
306
+ "dir": [math.cos(e) * math.cos(a), math.cos(e) * math.sin(a), math.sin(e)],
307
+ "zoom": float(zoom),
308
+ }
309
+
310
+
311
+ class coord_equal:
312
+ """Lock 2D axis units so shapes are not stretched to the panel.
313
+
314
+ ``ratio`` is the ggplot2 ``coord_fixed`` ratio: one unit on x has the
315
+ same on-screen length as ``ratio`` units on y. ``coord_equal()`` is
316
+ ``ratio=1``. The panel keeps its size; the camera shows extra range on
317
+ the looser axis instead of stretching the data.
318
+
319
+ Implicit-only figures (a circle, for example) use this automatically.
320
+ """
321
+
322
+ def __init__(self, ratio: float = 1.0):
323
+ try:
324
+ value = float(ratio)
325
+ except (TypeError, ValueError) as exc:
326
+ raise ValueError(
327
+ "coord_equal() ratio must be a positive number"
328
+ ) from exc
329
+ if not math.isfinite(value) or value <= 0.0:
330
+ raise ValueError("coord_equal() ratio must be a positive number")
331
+ self.ratio = value
332
+
333
+ def to_spec(self) -> dict:
334
+ return {"aspect": "equal", "ratio": self.ratio}
335
+
336
+
337
+ # ggplot2's coord_fixed(ratio=) is coord_equal by another name.
338
+ coord_fixed = coord_equal
339
+
340
+
341
+ class expand_limits:
342
+ """Make the axes reach these values even with no data there:
343
+ ``expand_limits(y=0)`` starts bars or a line at zero (ggplot2's
344
+ ``expand_limits``). Each argument is one value or a list."""
345
+
346
+ def __init__(self, x=None, y=None, colour=None, color=None):
347
+ def as_list(v):
348
+ if v is None:
349
+ return []
350
+ return list(v) if isinstance(v, (list, tuple)) else [v]
351
+
352
+ self.values = {"x": as_list(x), "y": as_list(y)}
353
+ if colour is not None or color is not None:
354
+ raise ValueError("expand_limits() takes x= and y=; use scale_colour_*(limits=) for colour")
355
+
356
+
357
+ class coord_polar:
358
+ """Draw ``r = f(theta)`` with equal units on x and y.
359
+
360
+ The angle runs from 0 to ``2π`` unless the layer sets ``tlim``.
361
+ A cardioid is ``geom_function("r = 1 + cos(theta)") + coord_polar()``.
362
+ """
363
+
364
+ def to_spec(self) -> dict:
365
+ return {"aspect": "equal", "ratio": 1.0}
366
+
367
+
368
+ class geom_surface(_Geom):
369
+ """3D surface from a regular x–y grid (height in ``z``).
370
+
371
+ Requires ``aes(x=, y=, z=)`` on a **complete rectangular grid** in long
372
+ form (one row per cell). Optional ``colour``/``fill`` colours vertices.
373
+ Forces 3D mode.
374
+
375
+ Parameters
376
+ ----------
377
+ wireframe:
378
+ If True, draw only mesh edges.
379
+ alpha:
380
+ Face opacity (default 0.95).
381
+ """
382
+
383
+ kind = "surface"
384
+
385
+ def __init__(self, mapping=None, *, wireframe: bool = False, alpha=0.95, **kw):
386
+ super().__init__(mapping, alpha=alpha, **kw)
387
+ self.wireframe = bool(wireframe)
388
+
389
+
390
+ class stat_density_3d:
391
+ """3D density-grid options for :class:`geom_isosurface`.
392
+
393
+ Not drawable alone. Add before ``geom_isosurface`` to set the histogram
394
+ resolution (and keep a ggplot2-shaped call site)::
395
+
396
+ ggplot(df, aes(x, y, z)) + stat_density_3d(n=24) + geom_isosurface(levels=[0.3, 0.7])
397
+ """
398
+
399
+ kind = "density_3d_stat"
400
+
401
+ def __init__(self, *, n: int = 32):
402
+ self.n = int(max(8, min(64, n)))
403
+
404
+
405
+ class geom_isosurface(_Geom):
406
+ """Isosurface of a 3D density estimate from point samples.
407
+
408
+ Requires ``aes(x=, y=, z=)`` on scatter-like data. v1 **embeds** density
409
+ estimation (histogram grid + light smoothing) then extracts surfaces at
410
+ the given levels. Levels are fractions of peak density in ``[0, 1]``.
411
+
412
+ Parameters
413
+ ----------
414
+ levels:
415
+ One or more relative thresholds (default ``[0.25, 0.5, 0.75]``).
416
+ n:
417
+ Density grid bins per axis (8–64). Overridden by a preceding
418
+ :class:`stat_density_3d` if present on the figure.
419
+ colour_by:
420
+ ``"level"`` colours mesh vertices by isolevel index (default).
421
+ wireframe, alpha:
422
+ Same idea as :class:`geom_surface`.
423
+ """
424
+
425
+ kind = "isosurface"
426
+
427
+ def __init__(
428
+ self,
429
+ mapping=None,
430
+ *,
431
+ levels: list[float] | tuple[float, ...] | None = None,
432
+ n: int = 32,
433
+ colour_by: str = "level",
434
+ wireframe: bool = False,
435
+ alpha: float = 0.55,
436
+ **kw,
437
+ ):
438
+ super().__init__(mapping, alpha=alpha, **kw)
439
+ if levels is None:
440
+ levels = (0.25, 0.5, 0.75)
441
+ self.levels = tuple(float(x) for x in levels)
442
+ if not self.levels:
443
+ raise ValueError("geom_isosurface() needs at least one level")
444
+ self.n = int(max(8, min(64, n)))
445
+ if colour_by not in {"level", "none"}:
446
+ raise ValueError("colour_by must be 'level' or 'none'")
447
+ self.colour_by = colour_by
448
+ self.wireframe = bool(wireframe)
449
+
450
+
451
+ class arrow:
452
+ """An arrowhead for geom_segment, geom_path, geom_line, annotate("segment").
453
+
454
+ ``angle`` in degrees, ``length`` in inches (ggplot2's 0.25), ``ends`` is
455
+ "last", "first", or "both", and ``type`` is "open" or "closed" (filled).
456
+ """
457
+
458
+ def __init__(self, angle=30.0, length=0.25, ends="last", type="open"): # noqa: A002
459
+ if ends not in {"last", "first", "both"}:
460
+ raise ValueError('arrow(ends=) is "last", "first", or "both"')
461
+ if type not in {"open", "closed"}:
462
+ raise ValueError('arrow(type=) is "open" or "closed"')
463
+ self.angle = float(angle)
464
+ self.length = float(length)
465
+ self.ends = ends
466
+ self.type = type
467
+
468
+ def spec(self) -> dict:
469
+ return {"angle": self.angle, "length": self.length * 96.0, "ends": self.ends, "type": self.type}
470
+
471
+
472
+ class geom_path(_Geom):
473
+ kind = "line"
474
+ sort_x = False # ggplot2 geom_path: connect in data order
475
+
476
+ def __init__(self, mapping=None, *, linewidth=None, width=None, linetype=None, arrow=None, **kw):
477
+ super().__init__(mapping, **kw)
478
+ self.arrow = arrow
479
+ self.linewidth = linewidth if linewidth is not None else (width or 2.0)
480
+ dash_pattern(linetype)
481
+ self.linetype = linetype
482
+
483
+
484
+ class geom_line(geom_path):
485
+ sort_x = True # ggplot2 geom_line: connect in order of x
486
+
487
+
488
+ _MARK_NAMES = ("roots", "extrema", "intersections")
489
+
490
+
491
+ def _mark_tuple(mark) -> tuple[str, ...]:
492
+ """``"roots"``, ``"roots, extrema"``, or a sequence of those names."""
493
+ if mark is None:
494
+ return ()
495
+ if isinstance(mark, str):
496
+ parts = [part.strip() for part in mark.split(",")]
497
+ parts = [part for part in parts if part]
498
+ elif isinstance(mark, (list, tuple)):
499
+ parts = []
500
+ for item in mark:
501
+ if not isinstance(item, str):
502
+ raise ValueError(
503
+ "mark must be 'roots', 'extrema', or 'intersections'"
504
+ )
505
+ text = item.strip()
506
+ if text:
507
+ parts.append(text)
508
+ else:
509
+ raise ValueError("mark must be 'roots', 'extrema', or 'intersections'")
510
+ if not parts or any(part not in _MARK_NAMES for part in parts):
511
+ raise ValueError("mark must be 'roots', 'extrema', or 'intersections'")
512
+ ordered: list[str] = []
513
+ for part in parts:
514
+ if part not in ordered:
515
+ ordered.append(part)
516
+ return tuple(ordered)
517
+
518
+
519
+ def _assignment_text(name: str, value) -> str | None:
520
+ """A formula string for this axis, or None when ``value`` is not text."""
521
+ if not isinstance(value, str):
522
+ return None
523
+ text = value.strip()
524
+ if not text:
525
+ raise ValueError(f"geom_function() {name}= is empty")
526
+ if "=" not in text:
527
+ text = f"{name} = {text}"
528
+ return text
529
+
530
+
531
+ class geom_function(_Geom):
532
+ """Draw a formula or callable as a curve or surface.
533
+
534
+ The quoted string is the form that works in scripts and notebooks::
535
+
536
+ ggplot() + geom_function("y = 2x + 2")
537
+ ggplot() + geom_function("z = sin(x) cos(y)", xlim=(-4, 4), ylim=(-4, 4))
538
+ ggplot() + geom_function("y = a x^2 + b", a=1, b=-2)
539
+ ggplot() + geom_function("y = a x^2") + transition_time(a=(0, 3))
540
+
541
+ A raw LaTeX string is accepted too. Without the ``r`` prefix Python
542
+ eats the backslashes (``\\frac`` becomes a form feed) before plot3
543
+ sees them::
544
+
545
+ ggplot() + geom_function(r"y = \\frac{\\sin x}{x}")
546
+
547
+ ``$...$`` or a backslash command selects LaTeX, and in that form
548
+ ``xy`` means ``x`` times ``y``. Plain text is unchanged, and braces
549
+ still group, so ``x^{2}`` works either way.
550
+
551
+ In a notebook, the same expression can be written without quotes
552
+ (``geom_function(y = 2*x + 2)``). A callable is full Python::
553
+
554
+ ggplot() + geom_function(lambda x: np.where(x < 0, 0, x**2))
555
+
556
+ The legend shows the formula typeset. A figure with one function puts
557
+ that formula in the title instead. ``label`` replaces it. ``$...$`` in
558
+ ``label`` or in ``labs()`` is math, and the rest of the string stays
559
+ plain::
560
+
561
+ geom_function("t = 2x^3 + 3y^3", label=r"Cubic: $t = 2x^3 + 3y^3$")
562
+
563
+ ``xlim`` / ``ylim`` / ``zlim`` are axis limits. On a curve, ``xlim`` is
564
+ the domain and ``ylim`` clips the view. On a surface, ``xlim`` and
565
+ ``ylim`` are the domain and ``zlim`` clips the view. ``n`` is the sample
566
+ count (default 501 on a curve, 80 per axis on a surface or implicit curve).
567
+ An implicit curve then subdivides the cells it crosses, so the line stays
568
+ smooth. A figure made only of implicit equations uses equal axis units,
569
+ so a circle stays round. A formula surface uses equal aspect (a cube)
570
+ unless you pass ``coord_3d``.
571
+
572
+ A parametric curve assigns two or three of x, y, and z, either in one
573
+ string or as keywords. ``tlim`` is the parameter interval (otherwise
574
+ ``xlim``, otherwise −10 to 10)::
575
+
576
+ geom_function("x = cos(t), y = sin(t)")
577
+ geom_function(x="cos(t)", y="sin(t)", z="t")
578
+
579
+ ``mark`` is ``"roots"``, ``"extrema"``, ``"intersections"``, or a
580
+ combination. ``area(0, 2)``, ``tangent(at=1)``, and ``derivative()``
581
+ attach to the curve that came before them. ``"y > x^2"`` shades the
582
+ side where the inequality holds. ``where(x < 0, 0, x^2)`` and a LaTeX
583
+ ``cases`` environment are piecewise.
584
+ """
585
+
586
+ # Formula coefficients (a=2, b=-3) arrive as keywords.
587
+ _takes_any_keyword = True
588
+
589
+ kind = "function"
590
+
591
+ def __init__(
592
+ self,
593
+ expr=None,
594
+ mapping=None,
595
+ *,
596
+ x=None,
597
+ y=None,
598
+ z=None,
599
+ f=None,
600
+ xlim=None,
601
+ ylim=None,
602
+ zlim=None,
603
+ tlim=None,
604
+ n=None,
605
+ linewidth=None,
606
+ wireframe: bool = False,
607
+ color=None,
608
+ colour=None,
609
+ alpha=None,
610
+ label=None,
611
+ mark=None,
612
+ **params,
613
+ ):
614
+ from plot3.expr import parse_formula
615
+
616
+ super().__init__(
617
+ mapping, color=color, colour=colour, alpha=alpha, **params
618
+ )
619
+ bound = dict(params)
620
+ pieces: list[str] = []
621
+ for name, value in (("x", x), ("y", y), ("z", z), ("f", f)):
622
+ piece = _assignment_text(name, value)
623
+ if piece is not None:
624
+ pieces.append(piece)
625
+ elif value is not None and expr is not None:
626
+ if isinstance(value, str):
627
+ raise ValueError(
628
+ "geom_function() takes one formula. "
629
+ 'Use geom_function(x="cos(t)", y="sin(t)") '
630
+ "with no positional formula, or pass numbers such as a=1."
631
+ )
632
+ bound[name] = value
633
+ elif value is not None and not isinstance(value, str):
634
+ bound[name] = value
635
+ if expr is not None and pieces:
636
+ raise ValueError(
637
+ "geom_function() takes one formula. "
638
+ 'Use geom_function(x="cos(t)", y="sin(t)") '
639
+ "with no positional formula, or pass numbers such as a=1."
640
+ )
641
+ if expr is not None:
642
+ formula = expr
643
+ elif len(pieces) >= 2:
644
+ formula = ", ".join(pieces)
645
+ elif len(pieces) == 1:
646
+ formula = pieces[0]
647
+ elif callable(f):
648
+ formula = f
649
+ else:
650
+ raise ValueError(
651
+ 'geom_function() needs a formula, for example '
652
+ 'geom_function("y = 2x + 2")'
653
+ )
654
+ # Defer unbound coefficients (``a`` in ``y = a x^2``). The transition
655
+ # is added with ``+`` afterwards, so it does not exist yet.
656
+ self.formula = parse_formula(formula, bound, defer_missing=True)
657
+ # Kept so a slider or transition that names a free symbol (t, mu)
658
+ # can re-read it as a coefficient at build time.
659
+ self._source = formula
660
+ self.xlim = xlim
661
+ self.ylim = ylim
662
+ self.zlim = zlim
663
+ self.tlim = tlim
664
+ self.n = None if n is None else int(n)
665
+ self.linewidth = None if linewidth is None else float(linewidth)
666
+ self.wireframe = bool(wireframe)
667
+ self.label = None if label is None else str(label)
668
+ self.marks = _mark_tuple(mark)
669
+ self.params = bound
670
+
671
+
672
+ class area:
673
+ """Shade ``y = f(x)`` from ``lo`` to ``hi`` and label the integral.
674
+
675
+ Add it after the curve. Limits swap when ``hi < lo``. ``baseline``
676
+ is the lower edge (default 0)::
677
+
678
+ ggplot() + geom_function("y = x^2") + area(0, 2)
679
+
680
+ Under a density the label is a probability, and a limit may be
681
+ infinite (it stops at the edge of the curve)::
682
+
683
+ geom_function("y = dnorm(x)") + area(-inf, -1.96) + area(1.96, inf)
684
+ """
685
+
686
+ def __init__(self, lo, hi, *, baseline=0):
687
+ try:
688
+ left = float(lo)
689
+ right = float(hi)
690
+ base = float(baseline)
691
+ except (TypeError, ValueError) as exc:
692
+ raise ValueError(
693
+ "area() needs numeric limits, for example area(0, 2)"
694
+ ) from exc
695
+ if any(math.isnan(value) for value in (left, right)) or not math.isfinite(base):
696
+ raise ValueError(
697
+ "area() needs numeric limits, for example area(0, 2)"
698
+ )
699
+ if left == right:
700
+ raise ValueError("area() needs two different limits")
701
+ if right < left:
702
+ left, right = right, left
703
+ self.lo = left
704
+ self.hi = right
705
+ self.baseline = base
706
+
707
+
708
+ class tangent:
709
+ """Tangent line to the preceding ``geom_function`` curve at ``x = at``."""
710
+
711
+ def __init__(self, at):
712
+ try:
713
+ value = float(at)
714
+ except (TypeError, ValueError) as exc:
715
+ raise ValueError(
716
+ "tangent() needs a number, for example tangent(at=1)"
717
+ ) from exc
718
+ if not math.isfinite(value):
719
+ raise ValueError(
720
+ "tangent() needs a number, for example tangent(at=1)"
721
+ )
722
+ self.at = value
723
+
724
+
725
+ class derivative:
726
+ """Draw ``f'`` of the preceding ``geom_function`` curve."""
727
+
728
+ def __init__(self):
729
+ return None
730
+
731
+
732
+ def _field_piece(name: str, value) -> str | None:
733
+ if value is None:
734
+ return None
735
+ if isinstance(value, str):
736
+ text = value.strip()
737
+ if not text:
738
+ raise ValueError(f"geom_vector_field() {name}= is empty")
739
+ if "=" not in text:
740
+ text = f"{name} = {text}"
741
+ return text
742
+ if isinstance(value, bool) or not isinstance(value, (int, float)):
743
+ raise ValueError(
744
+ f"geom_vector_field() {name}= must be a formula or a number"
745
+ )
746
+ number = float(value)
747
+ if not math.isfinite(number):
748
+ raise ValueError(
749
+ f"geom_vector_field() {name}= must be a formula or a number"
750
+ )
751
+ return f"{name} = {number}"
752
+
753
+
754
+ class geom_vector_field(_Geom):
755
+ """Arrows or streamlines for a planar field.
756
+
757
+ ``dx`` and ``dy`` are formulas in ``x`` and ``y``. The default window
758
+ is −2 to 2 with an 11 by 11 grid, so a rotation field stays readable.
759
+ ``stream=True`` draws unit-speed streamlines instead of arrows::
760
+
761
+ ggplot() + geom_vector_field("dx = -y, dy = x")
762
+ """
763
+
764
+ kind = "vector"
765
+
766
+ def __init__(
767
+ self,
768
+ expr=None,
769
+ mapping=None,
770
+ *,
771
+ dx=None,
772
+ dy=None,
773
+ dz=None,
774
+ xlim=(-2, 2),
775
+ ylim=(-2, 2),
776
+ zlim=None,
777
+ n=11,
778
+ stream: bool = False,
779
+ linewidth=None,
780
+ color=None,
781
+ colour=None,
782
+ alpha=None,
783
+ label=None,
784
+ **params,
785
+ ):
786
+ from plot3.expr import parse_formula
787
+
788
+ super().__init__(mapping, color=color, colour=colour, alpha=alpha)
789
+ if expr is not None and any(value is not None for value in (dx, dy, dz)):
790
+ raise ValueError(
791
+ 'geom_vector_field() takes one formula, for example '
792
+ 'geom_vector_field("dx = -y, dy = x")'
793
+ )
794
+ if expr is not None:
795
+ formula = expr
796
+ else:
797
+ parts = [
798
+ piece
799
+ for piece in (
800
+ _field_piece("dx", dx),
801
+ _field_piece("dy", dy),
802
+ _field_piece("dz", dz),
803
+ )
804
+ if piece is not None
805
+ ]
806
+ if len(parts) < 2:
807
+ raise ValueError(
808
+ 'geom_vector_field() needs dx and dy, for example '
809
+ '"dx = -y, dy = x"'
810
+ )
811
+ formula = ", ".join(parts)
812
+ self.formula = parse_formula(
813
+ formula, dict(params), defer_missing=True, role="field"
814
+ )
815
+ self.xlim = xlim
816
+ self.ylim = ylim
817
+ self.zlim = zlim
818
+ self.n = 11 if n is None else int(n)
819
+ self.stream = bool(stream)
820
+ self.linewidth = None if linewidth is None else float(linewidth)
821
+ self.label = None if label is None else str(label)
822
+ self.params = dict(params)
823
+
824
+
825
+ class geom_col(_Geom):
826
+ """Bars with heights from ``y`` (ggplot2 ``geom_col``).
827
+
828
+ Requires ``aes(x=, y=)``. ``x`` may be categorical or numeric.
829
+
830
+ Parameters
831
+ ----------
832
+ width:
833
+ Bar width as a fraction of the x resolution (ggplot2 default
834
+ ``0.9``). Data width is ``resolution(x) * width``. Override freely.
835
+ position:
836
+ With ``aes(colour=)`` groups: ``"stack"`` (default), ``"dodge"``
837
+ (side by side), ``"fill"`` (stacked to proportions), or
838
+ ``"identity"`` (overlapping). ``position_dodge(width=)`` also works.
839
+ """
840
+
841
+ kind = "col"
842
+
843
+ def __init__(self, mapping=None, *, width=0.9, position="stack", **kw):
844
+ super().__init__(mapping, **kw)
845
+ self.width = float(width)
846
+ self.position = position
847
+ self.data_override = None # optional layer-local frame (stats)
848
+
849
+
850
+ class geom_bar(_Geom):
851
+ """Count bars for a discrete ``x`` (ggplot2 ``geom_bar`` / ``stat_count``).
852
+
853
+ Only ``aes(x=)`` is required; counts become ``y``. Expanded to
854
+ ``geom_col`` at build time. Counted ``x`` values are drawn on a
855
+ **discrete** scale (like ``factor(x)`` in ggplot2). 2D only.
856
+
857
+ Parameters
858
+ ----------
859
+ width:
860
+ Bar width as a fraction of category spacing (ggplot2 default
861
+ ``0.9``). Use ``1.0`` for flush bars, smaller for more gap.
862
+ position:
863
+ With ``aes(colour=)`` groups the counts split by group:
864
+ ``"stack"`` (default), ``"dodge"``, ``"fill"``, or ``"identity"``.
865
+ """
866
+
867
+ kind = "bar"
868
+
869
+ def __init__(self, mapping=None, *, width=0.9, position="stack", **kw):
870
+ super().__init__(mapping, **kw)
871
+ self.width = float(width)
872
+ self.position = position
873
+
874
+
875
+ class position_dodge:
876
+ """Side by side within each x. ``width`` is the slot, as a fraction of x spacing."""
877
+
878
+ kind = "dodge"
879
+
880
+ def __init__(self, width=None):
881
+ self.width = None if width is None else float(width)
882
+
883
+
884
+ class position_dodge2:
885
+ """Side by side with a gap between neighbours (``padding``, a share of
886
+ each one's width), as ggplot2 dodges boxplots and bars."""
887
+
888
+ kind = "dodge"
889
+
890
+ def __init__(self, width=None, padding=0.1, preserve="total"):
891
+ self.width = None if width is None else float(width)
892
+ if not 0.0 <= float(padding) < 1.0:
893
+ raise ValueError("position_dodge2(padding=) is from 0 to less than 1")
894
+ self.padding = float(padding)
895
+ self.preserve = preserve
896
+
897
+
898
+ class position_jitter:
899
+ """Random offsets so overplotted points separate: ``geom_point(position=
900
+ position_jitter(width=0.2, height=0))``. ``None`` is 40% of the spacing
901
+ between values, as in ggplot2; ``seed`` keeps saved figures the same."""
902
+
903
+ kind = "jitter"
904
+
905
+ def __init__(self, width=None, height=None, seed=0):
906
+ self.width = width
907
+ self.height = height
908
+ self.seed = seed
909
+
910
+
911
+ class position_jitterdodge:
912
+ """Points beside each other by group within each x, then jittered: dots
913
+ over dodged boxplots. ``jitter_width`` defaults to 40% of the x spacing
914
+ shared among the groups; ``dodge_width`` matches geom_boxplot's 0.75."""
915
+
916
+ kind = "jitterdodge"
917
+
918
+ def __init__(self, jitter_width=None, jitter_height=0.0, dodge_width=0.75, seed=0):
919
+ self.jitter_width = jitter_width
920
+ self.jitter_height = float(jitter_height)
921
+ self.dodge_width = float(dodge_width)
922
+ self.seed = seed
923
+
924
+
925
+ class position_nudge:
926
+ """Shift by a fixed amount: labels just above their points with
927
+ ``geom_text(position=position_nudge(y=0.5))``."""
928
+
929
+ kind = "nudge"
930
+
931
+ def __init__(self, x=0.0, y=0.0):
932
+ self.x = float(x)
933
+ self.y = float(y)
934
+
935
+
936
+ class position_stack:
937
+ """Stacked, first group on top (ggplot2 order)."""
938
+
939
+ kind = "stack"
940
+ width = None
941
+
942
+
943
+ class position_fill:
944
+ """Stacked and scaled so each x sums to 1."""
945
+
946
+ kind = "fill"
947
+ width = None
948
+
949
+
950
+ class geom_jitter(_Geom):
951
+ """Points nudged at random so overplotted values separate (ggplot2 ``geom_jitter``).
952
+
953
+ ``width`` and ``height`` default to 40% of the spacing between x and y
954
+ values, as in ggplot2. Pass ``height=0`` to keep y exact. ``seed`` makes
955
+ the jitter repeatable, so a saved figure does not change between runs.
956
+ """
957
+
958
+ kind = "jitter"
959
+
960
+ def __init__(self, mapping=None, *, width=None, height=None, seed=0, size=None, **kw):
961
+ super().__init__(mapping, **kw)
962
+ self.width = width
963
+ self.height = height
964
+ self.seed = seed
965
+ self.size = size
966
+
967
+
968
+ class geom_crossbar(_Geom):
969
+ """A box from ``ymin`` to ``ymax`` with a line at ``y``: a mean and its
970
+ interval. Requires ``aes(x=, y=, ymin=, ymax=)``. ``fill`` colours the
971
+ box (its outline stays dark), ``colour`` the lines. ``fatten`` is how
972
+ much thicker the middle line is, as in ggplot2."""
973
+
974
+ kind = "crossbar"
975
+
976
+ def __init__(self, mapping=None, *, width=0.9, linewidth=1.0, fatten=2.5,
977
+ position="identity", **kw):
978
+ super().__init__(mapping, **kw)
979
+ self.width = float(width)
980
+ self.linewidth = float(linewidth)
981
+ self.fatten = float(fatten)
982
+ self.position = position
983
+
984
+
985
+ class geom_errorbarh(_Geom):
986
+ """Horizontal error bars from ``xmin`` to ``xmax`` at ``y``, with caps.
987
+ Requires ``aes(y=, xmin=, xmax=)``. ``height`` is the cap height as a
988
+ fraction of the y spacing."""
989
+
990
+ kind = "errorbarh"
991
+
992
+ def __init__(self, mapping=None, *, height=0.5, linewidth=1.0, **kw):
993
+ super().__init__(mapping, **kw)
994
+ self.height = float(height)
995
+ self.linewidth = float(linewidth)
996
+
997
+
998
+ class geom_polygon(_Geom):
999
+ """Filled polygons, one per ``group`` (or colour), corners in row order:
1000
+ maps, hulls, any closed shape. Concave shapes fill correctly."""
1001
+
1002
+ kind = "polygon"
1003
+
1004
+ def __init__(self, mapping=None, *, linewidth=0.5, **kw):
1005
+ super().__init__(mapping, **kw)
1006
+ self.linewidth = float(linewidth)
1007
+
1008
+
1009
+ class geom_blank(_Geom):
1010
+ """Draws nothing, but its data reach the scales: ``geom_blank(aes(y=0))``
1011
+ or a layer of limits from other data, as in ggplot2."""
1012
+
1013
+ kind = "blank"
1014
+
1015
+
1016
+ class geom_count(_Geom):
1017
+ """One point per distinct (x, y), its area by how many rows share it:
1018
+ overplotting made visible (ggplot2's geom_count). The size legend is
1019
+ titled ``n``."""
1020
+
1021
+ kind = "count"
1022
+
1023
+ def __init__(self, mapping=None, *, shape=None, **kw):
1024
+ super().__init__(mapping, **kw)
1025
+ self.shape = None if shape is None else shape_name(shape)
1026
+
1027
+
1028
+ class geom_bin_2d(_Geom):
1029
+ """Counts in rectangles over x and y, as a heatmap: big scatters made
1030
+ readable. ``bins`` (30) per axis, or ``binwidth``; a pair sets each axis."""
1031
+
1032
+ kind = "bin_2d"
1033
+
1034
+ def __init__(self, mapping=None, *, bins=30, binwidth=None, **kw):
1035
+ super().__init__(mapping, **kw)
1036
+ self.bins = bins
1037
+ self.binwidth = binwidth
1038
+
1039
+
1040
+ class geom_hex(_Geom):
1041
+ """Counts in hexagons over x and y (ggplot2's geom_hex)."""
1042
+
1043
+ kind = "hex"
1044
+
1045
+ def __init__(self, mapping=None, *, bins=30, binwidth=None, **kw):
1046
+ super().__init__(mapping, **kw)
1047
+ self.bins = bins
1048
+ self.binwidth = binwidth
1049
+
1050
+
1051
+ class geom_density_2d(_Geom):
1052
+ """Contour lines of a 2D kernel density (MASS::kde2d bandwidths), one set
1053
+ per colour group. ``bins`` (about 10 levels) or ``breaks`` sets them;
1054
+ ``n`` is the grid, ``h`` the bandwidths (as in kde2d)."""
1055
+
1056
+ kind = "density_2d"
1057
+
1058
+ def __init__(self, mapping=None, *, bins=None, breaks=None, n=100, h=None, linewidth=1.0, **kw):
1059
+ super().__init__(mapping, **kw)
1060
+ self.bins = bins
1061
+ self.breaks = breaks
1062
+ self.n = int(n)
1063
+ self.h = h
1064
+ self.linewidth = float(linewidth)
1065
+
1066
+
1067
+ class geom_density_2d_filled(geom_density_2d):
1068
+ """The 2D density in bands, one viridis colour per band, as ggplot2's
1069
+ geom_density_2d_filled."""
1070
+
1071
+ kind = "density_2d_filled"
1072
+
1073
+
1074
+ geom_density2d = geom_density_2d
1075
+ stat_density_2d = geom_density_2d
1076
+
1077
+
1078
+ class geom_contour(_Geom):
1079
+ """Contour lines of ``z`` on a regular grid of ``x`` and ``y``:
1080
+ ``aes(x=, y=, z=)``, one row per grid point. ``bins`` or ``breaks``
1081
+ sets the levels."""
1082
+
1083
+ kind = "contour"
1084
+
1085
+ def __init__(self, mapping=None, *, bins=None, breaks=None, linewidth=1.0, **kw):
1086
+ super().__init__(mapping, **kw)
1087
+ self.bins = bins
1088
+ self.breaks = breaks
1089
+ self.linewidth = float(linewidth)
1090
+
1091
+
1092
+ class stat_ellipse(_Geom):
1093
+ """A confidence ellipse per group: ``type="t"`` (robust, the default),
1094
+ ``"norm"``, or ``"euclid"``, at ``level`` (0.95)."""
1095
+
1096
+ kind = "ellipse"
1097
+
1098
+ def __init__(self, mapping=None, *, level=0.95, type="t", segments=51, linewidth=1.0, **kw): # noqa: A002
1099
+ super().__init__(mapping, **kw)
1100
+ if type not in {"t", "norm", "euclid"}:
1101
+ raise ValueError('stat_ellipse(type=) is "t", "norm", or "euclid"')
1102
+ if not 0 < float(level) < 1 and type != "euclid":
1103
+ raise ValueError("stat_ellipse(level=) is between 0 and 1")
1104
+ self.level = float(level)
1105
+ self.type = type
1106
+ self.segments = int(segments)
1107
+ self.linewidth = float(linewidth)
1108
+
1109
+
1110
+ class geom_freqpoly(_Geom):
1111
+ """A histogram drawn as a line through the bar tops (ggplot2
1112
+ ``geom_freqpoly``), one per colour group, ending at zero on both sides.
1113
+ Takes the histogram's ``bins``, ``binwidth``, and ``boundary``, and
1114
+ ``aes(y="after_stat(density)")`` for the density scale."""
1115
+
1116
+ kind = "freqpoly"
1117
+
1118
+ def __init__(self, mapping=None, *, bins=None, binwidth=None, method="fd",
1119
+ boundary=None, closed="right", linewidth=None, **kw):
1120
+ super().__init__(mapping, **kw)
1121
+ self.bins = bins
1122
+ self.binwidth = binwidth
1123
+ self.method = method
1124
+ self.boundary = boundary
1125
+ self.closed = closed
1126
+ self.linewidth = linewidth
1127
+
1128
+
1129
+ class geom_errorbar(_Geom):
1130
+ """Vertical error bars from ``ymin`` to ``ymax`` with caps.
1131
+
1132
+ Requires ``aes(x=, ymin=, ymax=)``. ``width`` is the cap width as a
1133
+ fraction of x spacing. ``position="dodge"`` lines up with dodged bars.
1134
+ """
1135
+
1136
+ kind = "errorbar"
1137
+
1138
+ def __init__(self, mapping=None, *, width=0.5, linewidth=1.0, position="identity", **kw):
1139
+ super().__init__(mapping, **kw)
1140
+ self.width = float(width)
1141
+ self.linewidth = float(linewidth)
1142
+ self.position = position
1143
+
1144
+
1145
+ class geom_linerange(_Geom):
1146
+ """A vertical line from ``ymin`` to ``ymax``. Requires ``aes(x=, ymin=, ymax=)``."""
1147
+
1148
+ kind = "linerange"
1149
+
1150
+ def __init__(self, mapping=None, *, linewidth=1.0, position="identity", **kw):
1151
+ super().__init__(mapping, **kw)
1152
+ self.linewidth = float(linewidth)
1153
+ self.position = position
1154
+
1155
+
1156
+ class geom_pointrange(_Geom):
1157
+ """A point at ``y`` on a line from ``ymin`` to ``ymax``.
1158
+
1159
+ Requires ``aes(x=, y=, ymin=, ymax=)``.
1160
+ """
1161
+
1162
+ kind = "pointrange"
1163
+
1164
+ def __init__(self, mapping=None, *, size=None, linewidth=1.0, position="identity", **kw):
1165
+ super().__init__(mapping, **kw)
1166
+ self.size = size
1167
+ self.linewidth = float(linewidth)
1168
+ self.position = position
1169
+
1170
+
1171
+ class geom_ribbon(_Geom):
1172
+ """A band between ``ymin`` and ``ymax`` along x (confidence bands, ranges).
1173
+
1174
+ Requires ``aes(x=, ymin=, ymax=)``; ``aes(colour=)`` draws one band per group.
1175
+ """
1176
+
1177
+ kind = "ribbon"
1178
+
1179
+
1180
+ class geom_smooth(_Geom):
1181
+ """A fitted trend with its confidence band (ggplot2 ``geom_smooth``).
1182
+
1183
+ ``method="loess"`` (default, local quadratic with ``span``) or ``"lm"``
1184
+ (a straight line). ``se=True`` shades the ``level`` confidence band,
1185
+ using Student's t like R. ``aes(colour=)`` fits each group.
1186
+ """
1187
+
1188
+ kind = "smooth"
1189
+
1190
+ def __init__(
1191
+ self, mapping=None, *, method="loess", se=True, level=0.95,
1192
+ span=0.75, n=80, linewidth=2.0, **kw,
1193
+ ):
1194
+ super().__init__(mapping, **kw)
1195
+ self.method = method
1196
+ self.se = bool(se)
1197
+ self.level = float(level)
1198
+ self.span = float(span)
1199
+ self.n = int(n)
1200
+ self.linewidth = float(linewidth)
1201
+
1202
+
1203
+ class stat_summary(_Geom):
1204
+ """Summarise ``y`` at each x, then draw it (ggplot2 ``stat_summary``).
1205
+
1206
+ ``fun_data`` is ``"mean_se"`` (default), ``"mean_cl_normal"`` (t
1207
+ interval), ``"mean_sdl"`` (mean ± 2 SD), ``"median_hilow"`` (median and
1208
+ the middle 95%), or a function returning ``(y, ymin, ymax)``.
1209
+ ``fun_args`` passes options, e.g. ``{"mult": 1}``. ``geom`` is
1210
+ ``"pointrange"`` (default), ``"errorbar"``, ``"linerange"``, ``"col"``,
1211
+ or ``"point"``.
1212
+ """
1213
+
1214
+ kind = "summary"
1215
+
1216
+ def __init__(
1217
+ self, mapping=None, *, fun_data="mean_se", fun_args=None, geom="pointrange",
1218
+ width=None, linewidth=1.0, size=None, position=None, **kw,
1219
+ ):
1220
+ super().__init__(mapping, **kw)
1221
+ self.fun_data = fun_data
1222
+ self.fun_args = dict(fun_args or {})
1223
+ self.geom = geom
1224
+ self.width = (0.9 if geom in {"col", "bar"} else 0.5) if width is None else float(width)
1225
+ self.linewidth = float(linewidth)
1226
+ self.size = size
1227
+ self.position = position
1228
+
1229
+
1230
+ # ggplot2's default shape palette, in order: what aes(shape=) assigns.
1231
+ SHAPE_ORDER = ["circle", "triangle", "square", "diamond", "plus", "cross"]
1232
+ _SHAPE_NUMBERS = {
1233
+ 0: "square", 1: "circle", 2: "triangle", 3: "plus", 4: "cross", 5: "diamond",
1234
+ 15: "square", 16: "circle", 17: "triangle", 18: "diamond", 19: "circle",
1235
+ 20: "circle", 21: "circle", 22: "square", 23: "diamond", 24: "triangle",
1236
+ }
1237
+ LINETYPE_ORDER = ["solid", "dashed", "dotted", "dotdash", "longdash", "twodash"]
1238
+
1239
+
1240
+ def shape_name(shape) -> str:
1241
+ """``"triangle"`` or R's numbers (17 = filled triangle) to a shape name."""
1242
+ if isinstance(shape, (int, float)) and not isinstance(shape, bool):
1243
+ if int(shape) in _SHAPE_NUMBERS:
1244
+ return _SHAPE_NUMBERS[int(shape)]
1245
+ name = str(shape).strip().lower()
1246
+ if name in SHAPE_ORDER:
1247
+ return name
1248
+ raise ValueError(f"shape {shape!r} is not one of {SHAPE_ORDER} or an R shape number")
1249
+
1250
+
1251
+ # ggplot2 line types as dash patterns, in multiples of the line width.
1252
+ LINETYPES = {
1253
+ "solid": None,
1254
+ "dashed": (4.0, 4.0),
1255
+ "dotted": (1.0, 3.0),
1256
+ "dotdash": (1.0, 3.0, 4.0, 3.0),
1257
+ "longdash": (8.0, 4.0),
1258
+ "twodash": (2.0, 2.0, 6.0, 2.0),
1259
+ }
1260
+ _LINETYPE_NUMBERS = ["blank", "solid", "dashed", "dotted", "dotdash", "longdash", "twodash"]
1261
+
1262
+
1263
+ def dash_pattern(linetype) -> tuple[float, ...] | None:
1264
+ """``"dashed"`` -> (4, 4). Also R's numbers (2 = dashed) and hex strings ("44")."""
1265
+ if linetype is None:
1266
+ return None
1267
+ if isinstance(linetype, (int, float)) and not isinstance(linetype, bool):
1268
+ index = int(linetype)
1269
+ if 0 <= index < len(_LINETYPE_NUMBERS):
1270
+ linetype = _LINETYPE_NUMBERS[index]
1271
+ name = str(linetype).strip().lower()
1272
+ if name in LINETYPES:
1273
+ return LINETYPES[name]
1274
+ if name and len(name) % 2 == 0 and all(c in "0123456789abcdef" for c in name):
1275
+ return tuple(float(int(c, 16)) for c in name)
1276
+ raise ValueError(
1277
+ f"linetype {linetype!r} is not one of {sorted(LINETYPES)} "
1278
+ "or a hex pattern such as '44'"
1279
+ )
1280
+
1281
+
1282
+ class geom_hline(_Geom):
1283
+ """Horizontal reference line(s) across the panel: ``geom_hline(yintercept=0)``.
1284
+
1285
+ ``yintercept`` may be a list. ``linetype`` is ``"solid"``, ``"dashed"``,
1286
+ ``"dotted"``, ``"dotdash"``, ``"longdash"``, or ``"twodash"``.
1287
+ """
1288
+
1289
+ kind = "hline"
1290
+
1291
+ def __init__(self, mapping=None, *, yintercept, linetype="solid", linewidth=1.0, **kw):
1292
+ super().__init__(mapping, **kw)
1293
+ self.values = _as_list(yintercept, "yintercept")
1294
+ self.linetype = linetype
1295
+ dash_pattern(linetype)
1296
+ self.linewidth = float(linewidth)
1297
+
1298
+
1299
+ class geom_vline(_Geom):
1300
+ """Vertical reference line(s): ``geom_vline(xintercept=[1, 2], linetype="dashed")``."""
1301
+
1302
+ kind = "vline"
1303
+
1304
+ def __init__(self, mapping=None, *, xintercept, linetype="solid", linewidth=1.0, **kw):
1305
+ super().__init__(mapping, **kw)
1306
+ self.values = _as_list(xintercept, "xintercept")
1307
+ self.linetype = linetype
1308
+ dash_pattern(linetype)
1309
+ self.linewidth = float(linewidth)
1310
+
1311
+
1312
+ class geom_abline(_Geom):
1313
+ """The line ``y = intercept + slope * x`` across the panel (default ``y = x``)."""
1314
+
1315
+ kind = "abline"
1316
+
1317
+ def __init__(self, mapping=None, *, slope=1.0, intercept=0.0, linetype="solid", linewidth=1.0, **kw):
1318
+ super().__init__(mapping, **kw)
1319
+ self.slope = float(slope)
1320
+ self.intercept = float(intercept)
1321
+ self.linetype = linetype
1322
+ dash_pattern(linetype)
1323
+ self.linewidth = float(linewidth)
1324
+
1325
+
1326
+ def _as_list(value, name: str) -> list:
1327
+ if value is None:
1328
+ raise ValueError(f"{name}= is required")
1329
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
1330
+ return list(value)
1331
+ return [value]
1332
+
1333
+
1334
+ class geom_text(_Geom):
1335
+ """Text at each row: ``geom_text(aes(label=name))`` (ggplot2 ``geom_text``).
1336
+
1337
+ ``size`` is in millimetres like ggplot2 (default 3.88, about 11 pt).
1338
+ ``hjust``/``vjust`` are 0 (left/bottom) to 1 (right/top). ``nudge_x`` and
1339
+ ``nudge_y`` shift labels off their points. ``check_overlap=True`` skips a
1340
+ label that would overlap one already drawn. ``fontface`` is ``"plain"``,
1341
+ ``"bold"``, ``"italic"``, or ``"bold.italic"``.
1342
+ """
1343
+
1344
+ kind = "text"
1345
+ _box = False
1346
+
1347
+ def __init__(
1348
+ self, mapping=None, *, size=3.88, hjust=0.5, vjust=0.5, nudge_x=0.0,
1349
+ nudge_y=0.0, check_overlap=False, fontface="plain", position="identity", **kw,
1350
+ ):
1351
+ super().__init__(mapping, **kw)
1352
+ self.size = float(size)
1353
+ self.hjust = float(hjust)
1354
+ self.vjust = float(vjust)
1355
+ self.nudge_x = float(nudge_x)
1356
+ self.nudge_y = float(nudge_y)
1357
+ if getattr(position, "kind", None) == "nudge":
1358
+ # position_nudge() is nudge_x / nudge_y, as in ggplot2.
1359
+ self.nudge_x += position.x
1360
+ self.nudge_y += position.y
1361
+ self.check_overlap = bool(check_overlap)
1362
+ if fontface not in {"plain", "bold", "italic", "bold.italic"}:
1363
+ raise ValueError("fontface is 'plain', 'bold', 'italic', or 'bold.italic'")
1364
+ self.fontface = fontface
1365
+
1366
+
1367
+ class geom_label(geom_text):
1368
+ """Like ``geom_text`` with a box behind each label (ggplot2 ``geom_label``)."""
1369
+
1370
+ kind = "text"
1371
+ _box = True
1372
+
1373
+
1374
+ def annotate(geom: str, *, x=None, y=None, xmin=None, xmax=None, ymin=None,
1375
+ ymax=None, xend=None, yend=None, label=None, **params):
1376
+ """One-off marks in data coordinates (ggplot2 ``annotate``).
1377
+
1378
+ * ``annotate("text", x=2, y=5, label="peak")`` (also ``"label"``)
1379
+ * ``annotate("rect", xmin=1, xmax=2, ymin=0, ymax=10, alpha=0.2)``
1380
+ * ``annotate("segment", x=1, y=1, xend=2, yend=3)``
1381
+ * ``annotate("point", x=1, y=1, size=8)``
1382
+
1383
+ Values may be lists for several marks. Text, segments, and points
1384
+ default to the ink colour; a rectangle to translucent grey.
1385
+ """
1386
+ def seq(value):
1387
+ if value is None:
1388
+ return None
1389
+ return list(value) if isinstance(value, (list, tuple, np.ndarray, pd.Series)) else [value]
1390
+
1391
+ def frame(**cols):
1392
+ given = {k: seq(v) for k, v in cols.items()}
1393
+ size = max(len(v) for v in given.values())
1394
+ return pd.DataFrame({k: v * size if len(v) == 1 else v for k, v in given.items()})
1395
+
1396
+ colour = params.pop("colour", params.pop("color", None))
1397
+ fill = params.pop("fill", None)
1398
+ kind = str(geom).lower()
1399
+ if kind in {"text", "label"}:
1400
+ if x is None or y is None or label is None:
1401
+ raise ValueError(f'annotate("{kind}") needs x=, y=, and label=')
1402
+ cls = geom_label if kind == "label" else geom_text
1403
+ out = cls(aes(x="x", y="y", label="label"), colour=colour, **params)
1404
+ out.data_override = frame(x=x, y=y, label=label)
1405
+ elif kind == "rect":
1406
+ if None in (xmin, xmax, ymin, ymax):
1407
+ raise ValueError('annotate("rect") needs xmin=, xmax=, ymin=, ymax=')
1408
+ box = frame(xmin=xmin, xmax=xmax, ymin=ymin, ymax=ymax)
1409
+ xs, ys, groups = [], [], []
1410
+ for row in box.itertuples(index=False):
1411
+ groups.append([len(xs), 4])
1412
+ xs += [row.xmin, row.xmin, row.xmax, row.xmax]
1413
+ ys += [row.ymin, row.ymax, row.ymax, row.ymin]
1414
+ out = _Geom(aes(x="x", y="y"), colour=fill or colour or "#7f7f7f",
1415
+ alpha=params.pop("alpha", 0.2))
1416
+ out.kind = "poly"
1417
+ out.data_override = pd.DataFrame({"x": xs, "y": ys})
1418
+ out._groups = groups
1419
+ out.linewidth = 0.0
1420
+ elif kind == "segment":
1421
+ if None in (x, y, xend, yend):
1422
+ raise ValueError('annotate("segment") needs x=, y=, xend=, yend=')
1423
+ seg = frame(x=x, y=y, xend=xend, yend=yend)
1424
+ xs, ys, groups = [], [], []
1425
+ for row in seg.itertuples(index=False):
1426
+ groups.append([len(xs), 2])
1427
+ xs += [row.x, row.xend]
1428
+ ys += [row.y, row.yend]
1429
+ out = geom_path(aes(x="x", y="y"), colour=colour, **params)
1430
+ out.data_override = pd.DataFrame({"x": xs, "y": ys})
1431
+ out._groups = groups
1432
+ out._ink_default = colour is None
1433
+ elif kind == "point":
1434
+ if x is None or y is None:
1435
+ raise ValueError('annotate("point") needs x= and y=')
1436
+ out = geom_point(aes(x="x", y="y"), colour=colour, **params)
1437
+ out.data_override = frame(x=x, y=y)
1438
+ out._ink_default = colour is None
1439
+ else:
1440
+ raise ValueError('annotate() draws "text", "label", "rect", "segment", or "point"')
1441
+ out._replace_mapping = True
1442
+ out._annotation = True
1443
+ return out
1444
+
1445
+
1446
+ class geom_tile(_Geom):
1447
+ """Heatmap cells: a rectangle at each (x, y), coloured by ``fill``.
1448
+
1449
+ x and y may be categories or numbers; ``width``/``height`` default to
1450
+ the spacing of the values. ``geom_raster`` is the same.
1451
+ """
1452
+
1453
+ kind = "tile"
1454
+
1455
+ def __init__(self, mapping=None, *, width=None, height=None, **kw):
1456
+ super().__init__(mapping, **kw)
1457
+ self.width = width
1458
+ self.height = height
1459
+
1460
+
1461
+ geom_raster = geom_tile
1462
+
1463
+
1464
+ class geom_area(_Geom):
1465
+ """Filled area under ``y``; groups (``fill``) stack, first level on top.
1466
+
1467
+ ``position="stack"`` (default), ``"fill"`` (shares of 1), or ``"identity"``.
1468
+ """
1469
+
1470
+ kind = "area_stat"
1471
+
1472
+ def __init__(self, mapping=None, *, position="stack", **kw):
1473
+ super().__init__(mapping, **kw)
1474
+ self.position = position
1475
+
1476
+
1477
+ class geom_step(_Geom):
1478
+ """A staircase line: ``direction="hv"`` (default), ``"vh"``, or ``"mid"``."""
1479
+
1480
+ kind = "step"
1481
+
1482
+ def __init__(self, mapping=None, *, direction="hv", linewidth=2.0, linetype=None, **kw):
1483
+ super().__init__(mapping, **kw)
1484
+ self.direction = direction
1485
+ self.linewidth = float(linewidth)
1486
+ dash_pattern(linetype)
1487
+ self.linetype = linetype
1488
+
1489
+
1490
+ class geom_segment(_Geom):
1491
+ """A line from (x, y) to (xend, yend) for every row."""
1492
+
1493
+ kind = "segment"
1494
+
1495
+ def __init__(self, mapping=None, *, linewidth=1.0, linetype=None, arrow=None, **kw):
1496
+ super().__init__(mapping, **kw)
1497
+ self.arrow = arrow
1498
+ self.linewidth = float(linewidth)
1499
+ dash_pattern(linetype)
1500
+ self.linetype = linetype
1501
+
1502
+
1503
+ class geom_rect(_Geom):
1504
+ """A rectangle from xmin..xmax and ymin..ymax for every row (shaded periods)."""
1505
+
1506
+ kind = "rect"
1507
+
1508
+
1509
+ class geom_qq(_Geom):
1510
+ """Q-Q plot: sorted ``aes(sample=)`` against normal quantiles."""
1511
+
1512
+ kind = "qq"
1513
+
1514
+ def __init__(self, mapping=None, *, size=None, **kw):
1515
+ super().__init__(mapping, **kw)
1516
+ self.size = size
1517
+
1518
+
1519
+ class geom_qq_line(_Geom):
1520
+ """The reference line of a Q-Q plot, through the quartiles (as in R)."""
1521
+
1522
+ kind = "qq_line"
1523
+
1524
+ def __init__(self, mapping=None, *, linewidth=1.5, linetype=None, **kw):
1525
+ super().__init__(mapping, **kw)
1526
+ self.linewidth = float(linewidth)
1527
+ dash_pattern(linetype)
1528
+ self.linetype = linetype
1529
+
1530
+
1531
+ stat_qq = geom_qq
1532
+ stat_qq_line = geom_qq_line
1533
+
1534
+
1535
+ class stat_ecdf(_Geom):
1536
+ """Empirical cumulative distribution of ``x`` as a step line, per group."""
1537
+
1538
+ kind = "ecdf"
1539
+
1540
+ def __init__(self, mapping=None, *, pad=True, linewidth=2.0, **kw):
1541
+ super().__init__(mapping, **kw)
1542
+ self.pad = bool(pad)
1543
+ self.linewidth = float(linewidth)
1544
+
1545
+
1546
+ class coord_flip:
1547
+ """Swap x and y: horizontal bars and boxplots, long category names on y."""
1548
+
1549
+ kind = "flip"
1550
+
1551
+
1552
+ class coord_cartesian:
1553
+ """Zoom to ``xlim`` / ``ylim`` without dropping data, as ggplot2 does.
1554
+
1555
+ ``scale_x_continuous(limits=)`` and ``xlim()`` remove the rows outside
1556
+ the limits before statistics run, so a boxplot or smoother changes.
1557
+ ``coord_cartesian`` computes everything from all rows and only moves the
1558
+ view. ``expand=False`` removes the small margin around the limits.
1559
+ """
1560
+
1561
+ kind = "cartesian"
1562
+
1563
+ def __init__(self, xlim=None, ylim=None, expand=True):
1564
+ for name, lim in (("xlim", xlim), ("ylim", ylim)):
1565
+ if lim is not None and (not isinstance(lim, (tuple, list)) or len(lim) != 2):
1566
+ raise ValueError(f"coord_cartesian({name}=) is a pair such as (0, 10)")
1567
+ self.xlim = None if xlim is None else tuple(xlim)
1568
+ self.ylim = None if ylim is None else tuple(ylim)
1569
+ self.expand = bool(expand)
1570
+
1571
+
1572
+ class geom_rug(_Geom):
1573
+ """Short ticks at the panel edges, one per row, showing where values fall.
1574
+
1575
+ ``sides`` is any of "b", "l", "t", "r" (default "bl": x along the
1576
+ bottom, y along the left). ``length`` is a fraction of the panel, as
1577
+ ggplot2's ``unit(0.03, "npc")``.
1578
+ """
1579
+
1580
+ kind = "rug"
1581
+
1582
+ def __init__(self, mapping=None, *, sides="bl", length=0.03, linewidth=0.5, **kw):
1583
+ super().__init__(mapping, **kw)
1584
+ sides = str(sides)
1585
+ if not sides or any(ch not in "bltr" for ch in sides):
1586
+ raise ValueError('geom_rug(sides=) uses "b", "l", "t", "r", for example "bl"')
1587
+ self.sides = sides
1588
+ self.length = float(length)
1589
+ self.linewidth = float(linewidth)
1590
+
1591
+
1592
+ class geom_histogram(_Geom):
1593
+ """Histogram of a continuous ``x`` (ggplot2 ``geom_histogram`` / ``stat_bin``).
1594
+
1595
+ Only ``aes(x=)`` is required. Bins are computed in Python and drawn as
1596
+ ``geom_col`` with **full bin width** so adjacent bars touch (no gaps).
1597
+ 2D only.
1598
+
1599
+ Parameters
1600
+ ----------
1601
+ bins:
1602
+ Explicit number of bins. When omitted (default), binning is chosen
1603
+ from the data via ``method`` (Freedman–Diaconis by default). Ignored
1604
+ when ``binwidth`` is set. Pass ``bins=30`` to force a fixed count
1605
+ (ggplot2's historical default).
1606
+ binwidth:
1607
+ Absolute bin width in data units. When set, overrides ``bins`` and
1608
+ ``method``.
1609
+ method:
1610
+ Automatic rule used only when both ``bins`` and ``binwidth`` are
1611
+ omitted: ``"fd"`` (Freedman–Diaconis, default), ``"scott"``,
1612
+ ``"sturges"``, ``"auto"`` (numpy's multi-rule choice), or other
1613
+ names accepted by ``numpy.histogram_bin_edges``.
1614
+ boundary:
1615
+ Optional bin boundary (ggplot2 ``boundary``). Aligns edges so that
1616
+ one edge falls on this value (modulo ``binwidth``).
1617
+ closed:
1618
+ ``"right"`` (default) or ``"left"`` — which side of each bin is
1619
+ closed (matches numpy / ggplot2 closed intervals).
1620
+ """
1621
+
1622
+ kind = "histogram"
1623
+
1624
+ def __init__(
1625
+ self,
1626
+ mapping=None,
1627
+ *,
1628
+ bins: int | None = None,
1629
+ binwidth: float | None = None,
1630
+ method: str = "fd",
1631
+ boundary: float | None = None,
1632
+ closed: str = "right",
1633
+ position="stack",
1634
+ **kw,
1635
+ ):
1636
+ super().__init__(mapping, **kw)
1637
+ # With aes(fill=g): "stack" (default), "dodge", "fill", "identity".
1638
+ self.position = position
1639
+ if bins is not None and int(bins) < 1:
1640
+ raise ValueError("bins must be positive")
1641
+ if binwidth is not None and float(binwidth) <= 0:
1642
+ raise ValueError("binwidth must be positive")
1643
+ if closed not in {"right", "left"}:
1644
+ raise ValueError("closed must be 'right' or 'left'")
1645
+ if not isinstance(method, str) or not method.strip():
1646
+ raise ValueError("method must be a non-empty string")
1647
+ self.bins = None if bins is None else int(bins)
1648
+ self.binwidth = None if binwidth is None else float(binwidth)
1649
+ self.method = method.strip().lower()
1650
+ self.boundary = None if boundary is None else float(boundary)
1651
+ self.closed = closed
1652
+ # Histograms use absolute bin width (bars touch); no relative width.
1653
+
1654
+
1655
+ class geom_boxplot(_Geom):
1656
+ """Box-and-whisker summary of ``y`` by ``x`` (ggplot2 ``geom_boxplot``).
1657
+
1658
+ Requires ``aes(x=, y=)``. ``x`` is usually categorical; ``y`` is numeric.
1659
+ Whiskers use the Tukey rule (``coef`` × IQR, default 1.5). Outliers are
1660
+ drawn as points. 2D only.
1661
+
1662
+ Parameters
1663
+ ----------
1664
+ width:
1665
+ Box width as a fraction of category spacing (default 0.75).
1666
+ outlier_size:
1667
+ Outlier point size in pixels (default 3).
1668
+ coef:
1669
+ Whisker fence multiplier on IQR (default 1.5). Set ``0`` to extend
1670
+ whiskers to the data min/max with no outliers.
1671
+ outliers:
1672
+ ``False`` draws no outlier points, as when jittered points already
1673
+ show every row. ``outlier_shape=None`` (ggplot2's
1674
+ ``outlier.shape = NA``) does the same.
1675
+ """
1676
+
1677
+ kind = "boxplot"
1678
+ _UNSET = object()
1679
+
1680
+ def __init__(
1681
+ self,
1682
+ mapping=None,
1683
+ *,
1684
+ width=0.75,
1685
+ outlier_size=3.0,
1686
+ coef=1.5,
1687
+ outliers=True,
1688
+ outlier_shape=_UNSET,
1689
+ position="dodge2",
1690
+ **kw,
1691
+ ):
1692
+ super().__init__(mapping, **kw)
1693
+ self.width = float(width)
1694
+ self.outlier_size = float(outlier_size)
1695
+ self.coef = float(coef)
1696
+ # Boxes grouped within an x sit side by side (ggplot2's dodge2).
1697
+ self.position = position
1698
+ self.outliers = bool(outliers) and outlier_shape is not None
1699
+
1700
+
1701
+ class geom_density(_Geom):
1702
+ """Kernel density estimate of a continuous variable (ggplot2 ``geom_density``).
1703
+
1704
+ Requires ``aes(x=)``. Optional ``colour``/``color`` draws one curve per
1705
+ group. Set ``fill=True`` (default) to shade under the curve. 2D only.
1706
+ """
1707
+
1708
+ kind = "density"
1709
+
1710
+ def __init__(
1711
+ self,
1712
+ mapping=None,
1713
+ *,
1714
+ n=512,
1715
+ adjust=1.0,
1716
+ fill=True,
1717
+ linewidth=1.5,
1718
+ **kw,
1719
+ ):
1720
+ super().__init__(mapping, **kw)
1721
+ self.n = int(n)
1722
+ self.adjust = float(adjust)
1723
+ self.fill = bool(fill)
1724
+ self.linewidth = float(linewidth)
1725
+
1726
+
1727
+ class geom_violin(_Geom):
1728
+ """Violin plot of ``y`` by ``x`` (ggplot2 ``geom_violin``).
1729
+
1730
+ Requires ``aes(x=, y=)``. Density is mirrored about each ``x`` category.
1731
+ 2D only.
1732
+ """
1733
+
1734
+ kind = "violin"
1735
+
1736
+ def __init__(
1737
+ self,
1738
+ mapping=None,
1739
+ *,
1740
+ n=128,
1741
+ adjust=1.0,
1742
+ width=0.9,
1743
+ linewidth=1.0,
1744
+ **kw,
1745
+ ):
1746
+ super().__init__(mapping, **kw)
1747
+ self.n = int(n)
1748
+ self.adjust = float(adjust)
1749
+ self.width = float(width)
1750
+ self.linewidth = float(linewidth)
1751
+
1752
+
1753
+ class facet_wrap:
1754
+ """Wrap panels by a discrete column (ggplot2 ``facet_wrap``).
1755
+
1756
+ Parameters
1757
+ ----------
1758
+ facets:
1759
+ Column name, or a formula-like string ``"~col"`` / ``". ~ col"``.
1760
+ ncol, nrow:
1761
+ Panel grid size. If both are omitted, the grid is ggplot2's: 3 panels
1762
+ in a row, 4 in 2 x 2, 5 or 6 in 2 rows of 3, 7 to 9 in 3 x 3.
1763
+ scales:
1764
+ ``"fixed"`` (shared domains across panels) or ``"free"`` (per-panel).
1765
+ """
1766
+
1767
+ def __init__(
1768
+ self,
1769
+ facets: str,
1770
+ *,
1771
+ ncol: int | None = None,
1772
+ nrow: int | None = None,
1773
+ scales: str = "fixed",
1774
+ labeller=None,
1775
+ ):
1776
+ if isinstance(facets, (list, tuple)):
1777
+ # facet_wrap(vars(cyl)): one column, as plot3 wraps by one.
1778
+ if len(facets) != 1:
1779
+ raise ValueError("facet_wrap() wraps by one column: facet_wrap(vars(cyl))")
1780
+ facets = facets[0]
1781
+ if not isinstance(facets, str) or not facets.strip():
1782
+ raise TypeError("facet_wrap() facets must be a column name string")
1783
+ name = facets.strip()
1784
+ if "~" in name:
1785
+ # Accept "~cyl", ". ~ cyl", "cyl ~ ."
1786
+ parts = [p.strip() for p in name.split("~")]
1787
+ candidates = [p for p in parts if p and p != "."]
1788
+ if len(candidates) != 1:
1789
+ raise ValueError(
1790
+ "facet_wrap() currently accepts a single facet column "
1791
+ f"(got {facets!r})"
1792
+ )
1793
+ name = candidates[0]
1794
+ if scales not in {"fixed", "free"}:
1795
+ raise ValueError("scales must be 'fixed' or 'free'")
1796
+ if ncol is not None and ncol < 1:
1797
+ raise ValueError("ncol must be positive")
1798
+ if nrow is not None and nrow < 1:
1799
+ raise ValueError("nrow must be positive")
1800
+ self.variable = name
1801
+ self.ncol = ncol
1802
+ self.nrow = nrow
1803
+ self.scales = scales
1804
+ self.labeller = _check_labeller(labeller)
1805
+
1806
+
1807
+ class facet_grid:
1808
+ """Panels in a grid: one row per level of ``rows``, one column per level
1809
+ of ``cols`` (ggplot2 ``facet_grid``).
1810
+
1811
+ ``facet_grid(rows="sex", cols="day")`` or the formula ``"sex ~ day"``
1812
+ (``". ~ day"`` for columns only). Column labels sit above the top row and
1813
+ row labels to the right, as in ggplot2. ``scales="fixed"`` (default)
1814
+ shares axes across panels; ``"free"`` lets each panel fit its data.
1815
+ """
1816
+
1817
+ def __init__(self, facets: str | None = None, *, rows=None, cols=None, scales: str = "fixed",
1818
+ labeller=None):
1819
+ if isinstance(facets, (list, tuple)):
1820
+ # facet_grid(vars(drv)) puts vars() in rows, as ggplot2 does.
1821
+ rows, facets = (facets if rows is None else rows), None
1822
+
1823
+ def one(value, side):
1824
+ if isinstance(value, (list, tuple)):
1825
+ if len(value) != 1:
1826
+ raise ValueError(f"facet_grid({side}=vars(...)) takes one column")
1827
+ return value[0]
1828
+ return value
1829
+
1830
+ rows, cols = one(rows, "rows"), one(cols, "cols")
1831
+ if facets is not None:
1832
+ if not isinstance(facets, str) or "~" not in facets:
1833
+ raise ValueError('facet_grid() takes "rows ~ cols", or rows= and cols=')
1834
+ left, right = (part.strip() for part in facets.split("~", 1))
1835
+ rows = rows if rows is not None else (left if left not in {"", "."} else None)
1836
+ cols = cols if cols is not None else (right if right not in {"", "."} else None)
1837
+ rows = _as_column_name(rows) if rows is not None else None
1838
+ cols = _as_column_name(cols) if cols is not None else None
1839
+ if rows is None and cols is None:
1840
+ raise ValueError("facet_grid() needs rows=, cols=, or both")
1841
+ if scales not in {"fixed", "free"}:
1842
+ raise ValueError("scales must be 'fixed' or 'free'")
1843
+ self.rows = rows
1844
+ self.cols = cols
1845
+ self.scales = scales
1846
+ self.labeller = _check_labeller(labeller)
1847
+
1848
+
1849
+ class _Vars(list):
1850
+ """Facet columns from vars()."""
1851
+
1852
+
1853
+ def vars(*names):
1854
+ """Facet columns, as ggplot2 writes them: ``facet_wrap(vars(cyl))``,
1855
+ ``facet_grid(rows=vars(drv), cols=vars(cyl))``.
1856
+
1857
+ ``from plot3 import *`` brings this ``vars`` in place of Python's, so
1858
+ Python's uses still work: ``vars(obj)`` returns its attributes and
1859
+ ``vars()`` the caller's local names.
1860
+ """
1861
+ import builtins
1862
+ import sys
1863
+
1864
+ if not names:
1865
+ return sys._getframe(1).f_locals
1866
+ if (
1867
+ len(names) == 1 and not isinstance(names[0], str)
1868
+ and _as_column_name(names[0]) is names[0] and hasattr(names[0], "__dict__")
1869
+ ):
1870
+ return builtins.vars(names[0])
1871
+ return _Vars(_as_column_name(n) for n in names)
1872
+
1873
+
1874
+ def label_value(variable, value) -> str:
1875
+ """A strip shows the level alone: ``high``."""
1876
+ return str(value)
1877
+
1878
+
1879
+ def label_both(variable, value) -> str:
1880
+ """A strip shows the column and the level: ``arm: high``."""
1881
+ return f"{variable}: {value}"
1882
+
1883
+
1884
+ def labeller(**by_variable):
1885
+ """A labeller per facet column: ``labeller(arm=label_both)``, or a dict
1886
+ of level names, ``labeller(sex={"F": "Female", "M": "Male"})``."""
1887
+ parts = {name: _check_labeller(rule) for name, rule in by_variable.items()}
1888
+
1889
+ def label(variable, value):
1890
+ rule = parts.get(variable)
1891
+ return strip_label(rule, variable, value)
1892
+
1893
+ return label
1894
+
1895
+
1896
+ def as_labeller(mapping):
1897
+ """Level names from a dict: ``as_labeller({"F": "Female", "M": "Male"})``."""
1898
+ return _check_labeller(dict(mapping))
1899
+
1900
+
1901
+ def _check_labeller(rule):
1902
+ if rule is None or callable(rule) or isinstance(rule, dict):
1903
+ return rule
1904
+ if isinstance(rule, str):
1905
+ named = {"label_value": label_value, "label_both": label_both}
1906
+ if rule not in named:
1907
+ raise ValueError('labeller= is "label_value", "label_both", a dict, or a function')
1908
+ return named[rule]
1909
+ raise TypeError('labeller= is "label_value", "label_both", a dict, or a function')
1910
+
1911
+
1912
+ def strip_label(rule, variable, value) -> str:
1913
+ """A facet strip's text under ``rule`` (a labeller, a dict, or None)."""
1914
+ if rule is None:
1915
+ return str(value)
1916
+ if isinstance(rule, dict):
1917
+ return str(rule.get(value, rule.get(str(value), value)))
1918
+ import inspect
1919
+
1920
+ try:
1921
+ params = [p for p in inspect.signature(rule).parameters.values()
1922
+ if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)]
1923
+ except (TypeError, ValueError):
1924
+ params = [None, None]
1925
+ return str(rule(value) if len(params) == 1 else rule(variable, value))
1926
+
1927
+
1928
+ _UNSET = object()
1929
+
1930
+
1931
+ class labs(dict):
1932
+ """Titles. ``colour`` (or ``fill``) names the colour legend.
1933
+
1934
+ ``subtitle`` sits under the title, ``caption`` at the bottom right (a
1935
+ data source, say), and ``tag`` at the top left ("A", "B" for the panels
1936
+ of a figure). ``None`` removes a title, as ggplot2's ``NULL`` does:
1937
+ ``labs(x=None)`` draws no x axis title.
1938
+ """
1939
+
1940
+ def __init__(self, title=_UNSET, x=_UNSET, y=_UNSET, z=_UNSET, color=_UNSET,
1941
+ colour=_UNSET, size=_UNSET, fill=_UNSET, subtitle=_UNSET,
1942
+ caption=_UNSET, tag=_UNSET, alpha=_UNSET):
1943
+ super().__init__()
1944
+ legend = color if color is not _UNSET else colour if colour is not _UNSET else fill
1945
+ for k, v in (("title", title), ("x", x), ("y", y), ("z", z),
1946
+ ("color", legend),
1947
+ ("size", size),
1948
+ ("subtitle", subtitle),
1949
+ ("caption", caption),
1950
+ ("tag", tag),
1951
+ ("alpha", alpha)):
1952
+ if v is not _UNSET:
1953
+ self[k] = "" if v is None else v
1954
+
1955
+
1956
+ def ggtitle(label, subtitle=None) -> labs:
1957
+ """The plot title (and subtitle): ``labs(title=, subtitle=)``."""
1958
+ return labs(title=label) if subtitle is None else labs(title=label, subtitle=subtitle)
1959
+
1960
+
1961
+ def xlab(label) -> labs:
1962
+ """The x axis title. ``xlab("")`` removes it."""
1963
+ return labs(x=label)
1964
+
1965
+
1966
+ def ylab(label) -> labs:
1967
+ """The y axis title. ``ylab("")`` removes it."""
1968
+ return labs(y=label)
1969
+
1970
+
1971
+ _GUIDE_KEYS = {"colour": "color", "color": "color", "fill": "color",
1972
+ "size": "size", "shape": "shape", "linetype": "linetype"}
1973
+
1974
+
1975
+ class guide_legend:
1976
+ """A legend of keys, with your ``title`` and order (``reverse=True``)."""
1977
+
1978
+ def __init__(self, title=None, reverse=False):
1979
+ self.title = title
1980
+ self.reverse = bool(reverse)
1981
+
1982
+
1983
+ class guide_colourbar(guide_legend):
1984
+ """A colour bar, with your ``title``."""
1985
+
1986
+
1987
+ guide_colorbar = guide_colourbar
1988
+
1989
+
1990
+ class guide_none:
1991
+ """No legend for this aesthetic."""
1992
+
1993
+
1994
+ class guides:
1995
+ """Hide or adjust a legend: ``guides(colour="none")``,
1996
+ ``guides(colour=guide_legend(title="Arm", reverse=True))``.
1997
+
1998
+ ``fill`` and ``colour`` share one legend in plot3, so either names it.
1999
+ ``"legend"`` and ``"colourbar"`` keep the default legend.
2000
+ """
2001
+
2002
+ def __init__(self, **kwargs):
2003
+ self.hidden: dict[str, bool] = {}
2004
+ self.options: dict[str, dict] = {}
2005
+ for key, value in kwargs.items():
2006
+ name = _GUIDE_KEYS.get(key)
2007
+ if name is None:
2008
+ raise ValueError(
2009
+ f"guides() takes colour, fill, size, shape, or linetype, not {key!r}"
2010
+ )
2011
+ if value is False or value is None or isinstance(value, guide_none) or (
2012
+ isinstance(value, str) and value == "none"
2013
+ ):
2014
+ self.hidden[name] = True
2015
+ elif isinstance(value, guide_legend):
2016
+ self.hidden[name] = False
2017
+ self.options[name] = {"title": value.title, "reverse": value.reverse}
2018
+ elif value in {"legend", "colourbar", "colorbar", True}:
2019
+ self.hidden[name] = False
2020
+ else:
2021
+ raise ValueError(
2022
+ f'guides({key}=) is "none", "legend", "colourbar", guide_legend(), '
2023
+ "or guide_colourbar()"
2024
+ )
2025
+
2026
+
2027
+ class scale_colour_continuous:
2028
+ """Numeric colour scale control.
2029
+
2030
+ trans: "linear" | "sqrt" | "log10"
2031
+ limits: (lo, hi) tuple, or "full" for the data min/max.
2032
+ Default (no scale added) is robust 2nd-98th percentile limits —
2033
+ skewed data (lidar intensity!) stays readable; values outside
2034
+ the limits clamp to the ramp ends.
2035
+ palette: "blue" (theme single-hue default) | "viridis" | "magma" | "turbo"
2036
+ """
2037
+
2038
+ def __init__(self, trans="linear", limits=None, palette="blue"):
2039
+ if trans not in ("linear", "sqrt", "log10"):
2040
+ raise ValueError("trans must be linear, sqrt or log10")
2041
+ if palette != "blue" and palette not in _CONT_PALETTES:
2042
+ raise ValueError(
2043
+ f"palette must be blue or one of {sorted(_CONT_PALETTES)}")
2044
+ self.trans = trans
2045
+ self.limits = limits
2046
+ self.palette = palette
2047
+
2048
+
2049
+ scale_color_continuous = scale_colour_continuous
2050
+
2051
+
2052
+ class scale_colour_viridis_c(scale_colour_continuous):
2053
+ """ggplot2-style viridis continuous scale: option viridis|magma|turbo."""
2054
+
2055
+ def __init__(self, option="viridis", trans="linear", limits=None):
2056
+ super().__init__(trans=trans, limits=limits, palette=option)
2057
+
2058
+
2059
+ scale_color_viridis_c = scale_colour_viridis_c
2060
+ scale_fill_viridis_c = scale_colour_viridis_c
2061
+
2062
+
2063
+ class scale_x_log10:
2064
+ """Base-10 logarithmic scale for x.
2065
+
2066
+ Positions are encoded in log10 space, so a decade is a constant distance
2067
+ and a tween along the axis moves in log space. Non-positive values are
2068
+ omitted. Tick labels stay in the original units (1, 10, 100, …).
2069
+ """
2070
+
2071
+ axis = "x"
2072
+
2073
+
2074
+ class scale_y_log10:
2075
+ """Base-10 logarithmic scale for y. See :class:`scale_x_log10`."""
2076
+
2077
+ axis = "y"
2078
+
2079
+
2080
+ def _as_range(value, name: str, who: str = "transition_time") -> tuple[float, float]:
2081
+ """A finite ``(lo, hi)`` pair. Inverted pairs are swapped."""
2082
+ pair = None
2083
+ if isinstance(value, (tuple, list, np.ndarray)) and not isinstance(
2084
+ value, (str, bytes)
2085
+ ):
2086
+ try:
2087
+ if len(value) == 2:
2088
+ pair = value
2089
+ except TypeError:
2090
+ pair = None
2091
+ if pair is not None:
2092
+ try:
2093
+ lo = float(pair[0])
2094
+ hi = float(pair[1])
2095
+ except (TypeError, ValueError):
2096
+ lo = hi = float("nan")
2097
+ else:
2098
+ if math.isfinite(lo) and math.isfinite(hi):
2099
+ if hi < lo:
2100
+ lo, hi = hi, lo
2101
+ return (lo, hi)
2102
+ raise TypeError(
2103
+ f"{who}() range {name} must be a pair of numbers, "
2104
+ f"for example {name}=(0, 3)"
2105
+ )
2106
+
2107
+
2108
+ class transition_time:
2109
+ """Animate points over a column, or a formula over parameter ranges.
2110
+
2111
+ Column form (Gapminder). Each ``aes(group=)`` value is one object.
2112
+ Rows are the keyframes of that object. The viewer stores every
2113
+ keyframe and interpolates in the browser, in scale space, with the
2114
+ axes held fixed on the range of every frame. The clock runs linearly
2115
+ from the first time value to the last, so a ten-year gap takes ten
2116
+ times as long as a one-year gap.
2117
+
2118
+ Range form. ``transition_time(a=(0, 3))`` sweeps every unbound
2119
+ coefficient of a ``geom_function`` together across ``frames`` steps
2120
+ (default 60). ``{frame_time}`` shows the parameters (``a = 1.50``).
2121
+
2122
+ (
2123
+ ggplot(df, aes(x="gdp", y="life", size="pop", colour="continent",
2124
+ group="country"))
2125
+ + geom_point()
2126
+ + scale_x_log10()
2127
+ + transition_time("year")
2128
+ + labs(title="{frame_time}")
2129
+ )
2130
+
2131
+ ggplot() + geom_function("y = a x^2") + transition_time(a=(0, 3))
2132
+ """
2133
+
2134
+ kind = "time"
2135
+
2136
+ def __init__(self, column=None, *, frames: int = 60, **ranges):
2137
+ if (column is None) == (not ranges):
2138
+ raise TypeError(
2139
+ "transition_time() takes a column (transition_time('year')) "
2140
+ "or parameter ranges (transition_time(a=(0, 3)))"
2141
+ )
2142
+ if column is None:
2143
+ self.column = None
2144
+ else:
2145
+ name = _as_column_name(column)
2146
+ if not isinstance(name, str) or not name:
2147
+ raise TypeError("transition_time() needs a column name")
2148
+ self.column = name
2149
+ self.ranges = {k: _as_range(v, k) for k, v in ranges.items()}
2150
+ nframes = int(frames)
2151
+ if self.ranges and nframes < 2:
2152
+ raise ValueError("transition_time() needs at least 2 frames")
2153
+ self.frames = nframes
2154
+
2155
+
2156
+ class slider:
2157
+ """Drag formula coefficients, one control per parameter.
2158
+
2159
+ Unlike :class:`transition_time`, each range moves on its own and nothing
2160
+ plays by itself. The grid is sampled in Python (``steps`` values along
2161
+ each range, last keyword varying fastest) and the viewer blends curves
2162
+ and surfaces between the neighboring cells. An implicit contour snaps
2163
+ to the nearest cell, because its vertex count changes.
2164
+
2165
+ ggplot() + geom_function("y = a sin(k x)") + slider(a=(0, 3), k=(1, 5))
2166
+
2167
+ ``steps`` defaults to 25 per parameter. A surface on the default grid
2168
+ is too large for several parameters at that count; pass a smaller
2169
+ ``steps`` or ``n``. Do not combine with ``transition_time()``.
2170
+ """
2171
+
2172
+ kind = "slider"
2173
+
2174
+ def __init__(self, *, steps: int = 25, **ranges):
2175
+ if not ranges:
2176
+ raise TypeError(
2177
+ "slider() needs a parameter range, for example slider(a=(0, 3))"
2178
+ )
2179
+ try:
2180
+ nsteps = int(steps)
2181
+ except (TypeError, ValueError):
2182
+ raise TypeError(
2183
+ "slider() steps= must be an integer, for example steps=25"
2184
+ ) from None
2185
+ if nsteps < 2:
2186
+ raise ValueError("slider() needs at least 2 steps")
2187
+ self.steps = nsteps
2188
+ self.ranges = {k: _as_range(v, k, "slider") for k, v in ranges.items()}
2189
+
2190
+
2191
+ class transition_states:
2192
+ """Animate a point layer over a discrete column.
2193
+
2194
+ Frames follow the order the states first appear in the data (so
2195
+ ``before`` then ``after`` stays in that order). Objects ease from one
2196
+ state to the next. ``aes(group=)`` matches the same object across states.
2197
+ ``{frame_time}`` shows the current state label.
2198
+ """
2199
+
2200
+ kind = "states"
2201
+
2202
+ def __init__(self, column):
2203
+ name = _as_column_name(column)
2204
+ if not isinstance(name, str) or not name:
2205
+ raise TypeError("transition_states() needs a column name")
2206
+ self.column = name
2207
+
2208
+
2209
+ class _Theme:
2210
+ """A named theme plus the font used by ``ggsave``.
2211
+
2212
+ ``base_size`` is in points, as in ggplot2. ``base_family`` is a CSS
2213
+ font-family list. Both apply when the figure is saved to PNG, SVG, or
2214
+ PDF. The interactive viewer keeps its own type.
2215
+ """
2216
+
2217
+ def __init__(
2218
+ self,
2219
+ name: str,
2220
+ *,
2221
+ base_size: float | None = None,
2222
+ base_family: str | None = None,
2223
+ ):
2224
+ if name not in _THEMES:
2225
+ known = ", ".join(sorted(_THEMES))
2226
+ raise ValueError(f"unknown theme {name!r}. Known themes: {known}")
2227
+ self.name = name
2228
+ self.base_size = _theme_points(base_size)
2229
+ self.base_family = None if base_family in (None, "") else str(base_family)
2230
+
2231
+
2232
+ def _theme_points(value) -> float | None:
2233
+ if value is None:
2234
+ return None
2235
+ if isinstance(value, bool):
2236
+ raise TypeError("base_size must be a font size in points, for example base_size=11")
2237
+ try:
2238
+ number = float(value)
2239
+ except (TypeError, ValueError) as exc:
2240
+ raise TypeError(
2241
+ "base_size must be a font size in points, for example base_size=11"
2242
+ ) from exc
2243
+ if not math.isfinite(number) or number < 1 or number > 96:
2244
+ raise ValueError("base_size must be between 1 and 96 points")
2245
+ return number
2246
+
2247
+
2248
+ def _theme(name: str, base_size=None, base_family=None) -> _Theme:
2249
+ return _Theme(name, base_size=base_size, base_family=base_family)
2250
+
2251
+
2252
+ def theme_dark(base_size=None, base_family=None) -> _Theme:
2253
+ return _theme("dark", base_size, base_family)
2254
+
2255
+
2256
+ def theme_light(base_size=None, base_family=None) -> _Theme:
2257
+ return _theme("light", base_size, base_family)
2258
+
2259
+
2260
+ def theme_bw(base_size=None, base_family=None) -> _Theme:
2261
+ """White page, grey grid, and a dark panel border."""
2262
+ return _theme("bw", base_size, base_family)
2263
+
2264
+
2265
+ def theme_classic(base_size=None, base_family=None) -> _Theme:
2266
+ """White page, no grid, and black axis lines."""
2267
+ return _theme("classic", base_size, base_family)
2268
+
2269
+
2270
+ def theme_minimal(base_size=None, base_family=None) -> _Theme:
2271
+ """White page, light grid, and no panel box."""
2272
+ return _theme("minimal", base_size, base_family)
2273
+
2274
+
2275
+ def theme_grey(base_size=None, base_family=None) -> _Theme:
2276
+ """ggplot2's default look: a grey panel with white grid lines."""
2277
+ return _theme("grey", base_size, base_family)
2278
+
2279
+
2280
+ theme_gray = theme_grey
2281
+
2282
+
2283
+ def theme_linedraw(base_size=None, base_family=None) -> _Theme:
2284
+ """A white panel with a thin dark grid and a black border."""
2285
+ return _theme("linedraw", base_size, base_family)
2286
+
2287
+
2288
+ def theme_lidar(base_size=None, base_family=None) -> _Theme:
2289
+ """A driving-scene look: black page, no box, grid, or ticks, points
2290
+ coloured by height from green through cyan to violet, and bright class
2291
+ colours for ``geom_box3d``."""
2292
+ return _theme("lidar", base_size, base_family)
2293
+
2294
+
2295
+ def theme_void(base_size=None, base_family=None) -> _Theme:
2296
+ """Only the data: no axes, ticks, grid, or panel box."""
2297
+ return _theme("void", base_size, base_family)
2298
+
2299
+
2300
+ class _ThemePatch:
2301
+ """Theme settings that do not change the colour theme."""
2302
+
2303
+ def __init__(self, legend_position=None, **options):
2304
+ self.legend_position = legend_position
2305
+ self.options = options
2306
+
2307
+
2308
+ class element_blank:
2309
+ """Draw nothing for this part: ``theme(panel_grid=element_blank())``."""
2310
+
2311
+
2312
+ class element_text:
2313
+ """Text settings for ``theme()``: ``element_text(angle=45, hjust=1)``."""
2314
+
2315
+ def __init__(self, size=None, colour=None, color=None, angle=None, hjust=None,
2316
+ vjust=None, face=None, family=None):
2317
+ self.size = size
2318
+ self.colour = colour if colour is not None else color
2319
+ self.angle = angle
2320
+ self.hjust = hjust
2321
+ self.vjust = vjust
2322
+ self.face = face
2323
+ self.family = family
2324
+
2325
+
2326
+ class element_line:
2327
+ """Line settings for ``theme()``: ``element_line(colour="grey80")``."""
2328
+
2329
+ def __init__(self, colour=None, color=None, linewidth=None, linetype=None, size=None):
2330
+ self.colour = colour if colour is not None else color
2331
+ self.linewidth = linewidth if linewidth is not None else size
2332
+ self.linetype = linetype
2333
+
2334
+
2335
+ class element_rect:
2336
+ """Box settings for ``theme()``: ``element_rect(fill="white", colour="black")``."""
2337
+
2338
+ def __init__(self, fill=None, colour=None, color=None, linewidth=None):
2339
+ self.fill = fill
2340
+ self.colour = colour if colour is not None else color
2341
+ self.linewidth = linewidth
2342
+
2343
+
2344
+ # Parts of a ggplot2 theme that plot3 does not draw separately: accepted so
2345
+ # R code ports, but they change nothing.
2346
+ _THEME_QUIET = {
2347
+ "panel_grid_minor", "panel_grid_minor_x", "panel_grid_minor_y", "axis_ticks",
2348
+ "axis_ticks_x", "axis_ticks_y", "axis_ticks_length", "legend_key",
2349
+ "legend_background", "plot_margin", "legend_key_size", "legend_text",
2350
+ "legend_box", "legend_justification", "legend_direction",
2351
+ }
2352
+
2353
+
2354
+ def _element_options(key: str, value, options: dict, tokens: dict) -> bool:
2355
+ """Fold one ggplot2 theme element into plot3's options and colours.
2356
+ False when plot3 has nothing that draws it."""
2357
+ from plot3.scaling import to_hex
2358
+
2359
+ blank = isinstance(value, element_blank)
2360
+ colour = getattr(value, "colour", None)
2361
+ fill = getattr(value, "fill", None)
2362
+ if key in {"axis_text_x", "axis_text_y", "axis_text"}:
2363
+ for axis in ("x", "y"):
2364
+ if key in {"axis_text", f"axis_text_{axis}"}:
2365
+ if blank:
2366
+ options[f"axis_text_{axis}"] = False
2367
+ if isinstance(value, element_text):
2368
+ if value.angle is not None and key in {"axis_text", "axis_text_x"}:
2369
+ angle = abs(float(value.angle))
2370
+ if not 0.0 <= angle <= 90.0:
2371
+ raise ValueError("axis text angles are from 0 to 90 degrees")
2372
+ options["axis_text_x_angle"] = angle
2373
+ if colour is not None:
2374
+ tokens["muted"] = to_hex(colour)
2375
+ return True
2376
+ if key in {"axis_title", "axis_title_x", "axis_title_y"}:
2377
+ for axis in ("x", "y"):
2378
+ if key in {"axis_title", f"axis_title_{axis}"} and blank:
2379
+ options[f"axis_title_{axis}"] = False
2380
+ if colour is not None:
2381
+ tokens["ink2"] = to_hex(colour)
2382
+ return True
2383
+ if key in {"panel_grid", "panel_grid_major", "panel_grid_major_x", "panel_grid_major_y"}:
2384
+ if blank:
2385
+ options["panel_grid"] = False
2386
+ elif colour is not None:
2387
+ tokens["grid"] = to_hex(colour)
2388
+ return True
2389
+ if key == "panel_background":
2390
+ if blank:
2391
+ tokens["panel"] = None
2392
+ elif fill is not None:
2393
+ tokens["panel"] = to_hex(fill)
2394
+ return True
2395
+ if key == "plot_background":
2396
+ if fill is not None:
2397
+ tokens["surface"] = to_hex(fill)
2398
+ return True
2399
+ if key == "panel_border":
2400
+ if blank:
2401
+ tokens["frame"] = "none"
2402
+ else:
2403
+ tokens["frame"] = "box"
2404
+ if colour is not None:
2405
+ tokens["axis"] = to_hex(colour)
2406
+ return True
2407
+ if key in {"axis_line", "axis_line_x", "axis_line_y"}:
2408
+ if not blank:
2409
+ tokens["frame"] = "axes"
2410
+ if colour is not None:
2411
+ tokens["axis"] = to_hex(colour)
2412
+ return True
2413
+ if key == "legend_title":
2414
+ if blank:
2415
+ options["legend_title"] = False
2416
+ return True
2417
+ if key == "plot_title":
2418
+ if isinstance(value, element_text):
2419
+ if value.hjust is not None:
2420
+ options["plot_title_hjust"] = float(value.hjust)
2421
+ if colour is not None:
2422
+ tokens["ink"] = to_hex(colour)
2423
+ return True
2424
+ if key == "text":
2425
+ if isinstance(value, element_text):
2426
+ if value.size is not None:
2427
+ options["base_size"] = float(value.size)
2428
+ if value.family is not None:
2429
+ options["base_family"] = str(value.family)
2430
+ if colour is not None:
2431
+ tokens.update({"ink": to_hex(colour), "ink2": to_hex(colour), "muted": to_hex(colour)})
2432
+ return True
2433
+ return key in _THEME_QUIET
2434
+
2435
+
2436
+ def theme(
2437
+ *,
2438
+ legend_position=None,
2439
+ legend_title=None,
2440
+ panel_grid=None,
2441
+ axis_text_x_angle=None,
2442
+ plot_title_hjust=None,
2443
+ base_size=None,
2444
+ base_family=None,
2445
+ **elements,
2446
+ ) -> _ThemePatch:
2447
+ """Change parts of the theme; later ``theme()`` calls add to earlier ones.
2448
+
2449
+ ``legend_position``: ``"right"``, ``"bottom"``, ``"none"``, or a pair
2450
+ ``(x, y)`` in 0–1 panel coordinates (inside the panel).
2451
+ ``legend_title``: ``False`` hides the legend title.
2452
+ ``panel_grid``: ``False`` removes the grid lines.
2453
+ ``axis_text_x_angle``: ``45`` or ``90`` turns long x labels.
2454
+ ``plot_title_hjust``: ``0`` left (default), ``0.5`` centred, ``1`` right.
2455
+ ``base_size`` (points) and ``base_family`` set the type for saved files.
2456
+
2457
+ ggplot2's elements work too, with dots or underscores:
2458
+ ``theme(axis_text_x=element_text(angle=45))``,
2459
+ ``theme(**{"panel.grid": element_blank()})``,
2460
+ ``panel_background=element_rect(fill="grey95")``,
2461
+ ``axis_title_y=element_blank()``, ``plot_title=element_text(hjust=0.5)``.
2462
+ """
2463
+ import warnings
2464
+
2465
+ options = {}
2466
+ tokens: dict = {}
2467
+ if legend_title is not None:
2468
+ options["legend_title"] = bool(legend_title)
2469
+ if isinstance(panel_grid, element_blank):
2470
+ options["panel_grid"] = False
2471
+ elif panel_grid is not None and not isinstance(panel_grid, (element_line, element_rect)):
2472
+ options["panel_grid"] = bool(panel_grid)
2473
+ elif panel_grid is not None:
2474
+ _element_options("panel_grid", panel_grid, options, tokens)
2475
+ if axis_text_x_angle is not None:
2476
+ angle = float(axis_text_x_angle)
2477
+ if not 0.0 <= angle <= 90.0:
2478
+ raise ValueError("axis_text_x_angle is from 0 to 90 degrees")
2479
+ options["axis_text_x_angle"] = angle
2480
+ if plot_title_hjust is not None:
2481
+ hjust = float(plot_title_hjust)
2482
+ if not 0.0 <= hjust <= 1.0:
2483
+ raise ValueError("plot_title_hjust is from 0 (left) to 1 (right)")
2484
+ options["plot_title_hjust"] = hjust
2485
+ if base_size is not None:
2486
+ options["base_size"] = float(base_size)
2487
+ if base_family is not None:
2488
+ options["base_family"] = str(base_family)
2489
+ for raw, value in elements.items():
2490
+ key = raw.replace(".", "_")
2491
+ if key == "legend_position":
2492
+ legend_position = value
2493
+ continue
2494
+ if not isinstance(value, (element_blank, element_text, element_line, element_rect)):
2495
+ raise TypeError(
2496
+ f"theme({raw}=) takes element_text(), element_line(), element_rect(), "
2497
+ "or element_blank()"
2498
+ )
2499
+ if not _element_options(key, value, options, tokens):
2500
+ warnings.warn(f"theme(): plot3 does not draw {raw}; it is ignored", stacklevel=2)
2501
+ if tokens:
2502
+ options["tokens"] = tokens
2503
+ return _ThemePatch(_check_legend_position(legend_position), **options)
2504
+
2505
+
2506
+ def _check_legend_position(value):
2507
+ if value is None:
2508
+ return None
2509
+ if isinstance(value, str):
2510
+ name = value.strip().lower()
2511
+ if name not in {"right", "bottom", "none"}:
2512
+ raise ValueError(
2513
+ "legend_position must be 'right', 'bottom', 'none', "
2514
+ "or a pair (x, y) from 0 to 1"
2515
+ )
2516
+ return name
2517
+ if isinstance(value, (tuple, list)) and len(value) == 2:
2518
+ try:
2519
+ x = float(value[0])
2520
+ y = float(value[1])
2521
+ except (TypeError, ValueError) as exc:
2522
+ raise ValueError(
2523
+ "legend_position (x, y) uses numbers from 0 to 1"
2524
+ ) from exc
2525
+ if not (
2526
+ math.isfinite(x) and math.isfinite(y) and 0.0 <= x <= 1.0 and 0.0 <= y <= 1.0
2527
+ ):
2528
+ raise ValueError("legend_position (x, y) uses numbers from 0 to 1")
2529
+ return (x, y)
2530
+ raise ValueError(
2531
+ "legend_position must be 'right', 'bottom', 'none', or a pair (x, y)"
2532
+ )
2533
+
2534
+
2535
+ # ggplot2's stat_* spellings of the same layers.
2536
+ stat_smooth = geom_smooth
2537
+ stat_bin = geom_histogram
2538
+ stat_count = geom_bar
2539
+ stat_density = geom_density
2540
+
2541
+
2542
+ def stat_function(mapping=None, *, fun=None, args=None, **kwargs):
2543
+ """ggplot2's ``stat_function(fun=dnorm, args={"mean": 2})``: a curve of
2544
+ ``fun(x, **args)``. ``fun`` may also be a formula string."""
2545
+ if fun is None:
2546
+ raise ValueError('stat_function() needs fun=, for example fun=dnorm or fun="sin(x)"')
2547
+ if args and callable(fun):
2548
+ bound = dict(args)
2549
+ name = getattr(fun, "__name__", "f")
2550
+
2551
+ def curve(x, _f=fun, _a=bound):
2552
+ return _f(x, **_a)
2553
+
2554
+ curve.__name__ = name
2555
+ return geom_function(curve, mapping, **kwargs)
2556
+ if args:
2557
+ return geom_function(fun, mapping, **dict(args), **kwargs)
2558
+ return geom_function(fun, mapping, **kwargs)