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/scaling.py ADDED
@@ -0,0 +1,636 @@
1
+ """ggplot2's scale functions: axis limits, breaks, labels, and palettes.
2
+
3
+ Position::
4
+
5
+ scale_x_continuous(name="Dose (mg)", limits=(0, 10), breaks=[0, 5, 10])
6
+ scale_y_continuous(labels="percent") # 0.25 -> 25%
7
+ scale_x_discrete(limits=["low", "mid", "high"], labels={"mid": "medium"})
8
+ scale_y_reverse(), xlim(0, 10), ylim("a", "b"), lims(x=(0, 1))
9
+ scale_x_date(date_breaks="3 months", date_labels="%b %Y")
10
+
11
+ Colour and fill (one channel: fill colours filled shapes, colour the rest)::
12
+
13
+ scale_colour_manual(values={"ctrl": "grey", "drug": "firebrick"})
14
+ scale_fill_brewer(palette="Set2"), scale_colour_viridis_d()
15
+ scale_colour_gradient(low="white", high="darkblue")
16
+ scale_fill_gradient2(low="blue", mid="white", high="red", midpoint=0)
17
+
18
+ Shape and linetype::
19
+
20
+ scale_shape_manual(values=["circle", "triangle"])
21
+ scale_linetype_manual(values=["solid", "dashed"])
22
+
23
+ Limits on a continuous axis drop rows outside them, as in ggplot2 (and say
24
+ so). A function's samples and computed layers are clipped by the panel.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import colorsys
30
+ import math
31
+ from typing import Any
32
+
33
+ import numpy as np
34
+
35
+ # ── label formats ────────────────────────────────────────────────────────────
36
+
37
+
38
+ def _trim(number: float, digits: int = 6) -> str:
39
+ text = f"{number:.{digits}f}".rstrip("0").rstrip(".")
40
+ return "0" if text in {"-0", ""} else text
41
+
42
+
43
+ def format_label(value: float, labels: Any) -> str:
44
+ """One tick label: "percent", "comma", "dollar", "scientific", a format
45
+ string such as "{:.1f} kg", or a function of the value."""
46
+ if callable(labels):
47
+ return str(labels(value))
48
+ if labels == "percent":
49
+ return f"{_trim(value * 100.0, 4)}%"
50
+ if labels == "comma":
51
+ return f"{value:,.0f}" if float(value).is_integer() else f"{value:,.2f}"
52
+ if labels == "dollar":
53
+ sign = "-" if value < 0 else ""
54
+ body = f"{abs(value):,.0f}" if float(value).is_integer() else f"{abs(value):,.2f}"
55
+ return f"{sign}${body}"
56
+ if labels == "scientific":
57
+ return f"{value:.2e}"
58
+ if isinstance(labels, str) and "{" in labels:
59
+ return labels.format(value)
60
+ raise ValueError(
61
+ 'labels is "percent", "comma", "dollar", "scientific", a format string '
62
+ 'such as "{:.1f}", a function, or a list matching breaks'
63
+ )
64
+
65
+
66
+ # ── position scales ──────────────────────────────────────────────────────────
67
+
68
+ _DATE_UNITS = {
69
+ "sec": "s", "second": "s", "min": "min", "minute": "min", "hour": "h",
70
+ "day": "D", "week": "W-MON", "month": "MS", "quarter": "QS", "year": "YS",
71
+ }
72
+
73
+
74
+ def date_freq(spec: str) -> str:
75
+ """ "3 months" -> "3MS" (a pandas frequency)."""
76
+ parts = str(spec).strip().lower().split()
77
+ count, unit = (parts[0], parts[1]) if len(parts) == 2 else ("1", parts[0])
78
+ unit = unit.rstrip("s") if unit not in {"s"} else unit
79
+ if unit not in _DATE_UNITS:
80
+ raise ValueError(
81
+ f"date_breaks {spec!r}: use a count and a unit, such as '3 months', "
82
+ f"'1 week', '2 years' ({', '.join(sorted(_DATE_UNITS))})"
83
+ )
84
+ return f"{int(count)}{_DATE_UNITS[unit]}"
85
+
86
+
87
+ class PositionScale:
88
+ """One x or y scale. ``kind`` is continuous, discrete, or date."""
89
+
90
+ def __init__(self, axis, kind, *, name=None, limits=None, breaks=None,
91
+ labels=None, trans=None, date_breaks=None, date_labels=None):
92
+ self.axis = axis
93
+ self.kind = kind
94
+ self.name = name
95
+ self.limits = limits
96
+ self.breaks = None if breaks is None else list(breaks)
97
+ self.labels = labels
98
+ self.trans = trans
99
+ self.date_breaks = date_breaks
100
+ self.date_labels = date_labels
101
+ if kind == "continuous" and limits is not None:
102
+ if len(limits) != 2:
103
+ raise ValueError(f"{axis} limits are (low, high); use None for an open end")
104
+ if trans not in {None, "log10", "reverse", "identity"}:
105
+ raise ValueError("trans is 'log10', 'reverse', or None")
106
+ if isinstance(labels, (list, tuple)) and self.breaks is not None and len(labels) != len(self.breaks):
107
+ raise ValueError("labels must have one entry per break")
108
+ if date_breaks is not None:
109
+ date_freq(date_breaks)
110
+
111
+
112
+ def _scale(axis, kind, **kw):
113
+ return PositionScale(axis, kind, **kw)
114
+
115
+
116
+ def scale_x_continuous(name=None, *, limits=None, breaks=None, labels=None, trans=None):
117
+ """x axis title, ``limits=(lo, hi)``, ``breaks=[...]``, ``labels=``, ``trans=``."""
118
+ return _scale("x", "continuous", name=name, limits=limits, breaks=breaks, labels=labels, trans=trans)
119
+
120
+
121
+ def scale_y_continuous(name=None, *, limits=None, breaks=None, labels=None, trans=None):
122
+ """y axis: see :func:`scale_x_continuous`."""
123
+ return _scale("y", "continuous", name=name, limits=limits, breaks=breaks, labels=labels, trans=trans)
124
+
125
+
126
+ def scale_x_reverse(name=None, *, limits=None, breaks=None, labels=None):
127
+ """x runs from high to low."""
128
+ return _scale("x", "continuous", name=name, limits=limits, breaks=breaks, labels=labels, trans="reverse")
129
+
130
+
131
+ def scale_y_reverse(name=None, *, limits=None, breaks=None, labels=None):
132
+ """y runs from high to low (depth, rank)."""
133
+ return _scale("y", "continuous", name=name, limits=limits, breaks=breaks, labels=labels, trans="reverse")
134
+
135
+
136
+ def scale_x_discrete(name=None, *, limits=None, labels=None):
137
+ """Category order (``limits``, which also drops the rest) and display ``labels``."""
138
+ return _scale("x", "discrete", name=name, limits=None if limits is None else [str(v) for v in limits], labels=labels)
139
+
140
+
141
+ def scale_y_discrete(name=None, *, limits=None, labels=None):
142
+ return _scale("y", "discrete", name=name, limits=None if limits is None else [str(v) for v in limits], labels=labels)
143
+
144
+
145
+ def scale_x_date(name=None, *, limits=None, date_breaks=None, date_labels=None):
146
+ """Dates on x: ``date_breaks="3 months"``, ``date_labels="%b %Y"``."""
147
+ return _scale("x", "date", name=name, limits=limits, date_breaks=date_breaks, date_labels=date_labels)
148
+
149
+
150
+ def scale_y_date(name=None, *, limits=None, date_breaks=None, date_labels=None):
151
+ return _scale("y", "date", name=name, limits=limits, date_breaks=date_breaks, date_labels=date_labels)
152
+
153
+
154
+ scale_x_datetime = scale_x_date
155
+ scale_y_datetime = scale_y_date
156
+
157
+
158
+ def _lim(axis, values):
159
+ if len(values) == 1 and isinstance(values[0], (list, tuple)):
160
+ values = tuple(values[0])
161
+ if len(values) == 2 and all(v is None or isinstance(v, (int, float, np.integer, np.floating)) for v in values):
162
+ return _scale(axis, "continuous", limits=tuple(values))
163
+ if all(isinstance(v, str) for v in values):
164
+ return _scale(axis, "discrete", limits=[str(v) for v in values])
165
+ if len(values) == 2:
166
+ return _scale(axis, "date", limits=tuple(values))
167
+ raise ValueError(f"{axis}lim() takes two numbers, two dates, or category names")
168
+
169
+
170
+ def xlim(*values):
171
+ """``xlim(0, 10)`` (rows outside are dropped) or ``xlim("a", "b")`` (order)."""
172
+ return _lim("x", values)
173
+
174
+
175
+ def ylim(*values):
176
+ """``ylim(0, 100)`` or ``ylim("low", "high")``."""
177
+ return _lim("y", values)
178
+
179
+
180
+ class _Lims(list):
181
+ """Several scales at once."""
182
+
183
+
184
+ def lims(*, x=None, y=None):
185
+ out = _Lims()
186
+ if x is not None:
187
+ out.append(xlim(x))
188
+ if y is not None:
189
+ out.append(ylim(y))
190
+ return out
191
+
192
+
193
+ # ── colour and fill ──────────────────────────────────────────────────────────
194
+
195
+ _BREWER = {
196
+ "Set1": ["#E41A1C", "#377EB8", "#4DAF4A", "#984EA3", "#FF7F00", "#FFFF33", "#A65628", "#F781BF", "#999999"],
197
+ "Set2": ["#66C2A5", "#FC8D62", "#8DA0CB", "#E78AC3", "#A6D854", "#FFD92F", "#E5C494", "#B3B3B3"],
198
+ "Set3": ["#8DD3C7", "#FFFFB3", "#BEBADA", "#FB8072", "#80B1D3", "#FDB462", "#B3DE69", "#FCCDE5", "#D9D9D9", "#BC80BD", "#CCEBC5", "#FFED6F"],
199
+ "Dark2": ["#1B9E77", "#D95F02", "#7570B3", "#E7298A", "#66A61E", "#E6AB02", "#A6761D", "#666666"],
200
+ "Paired": ["#A6CEE3", "#1F78B4", "#B2DF8A", "#33A02C", "#FB9A99", "#E31A1C", "#FDBF6F", "#FF7F00", "#CAB2D6", "#6A3D9A", "#FFFF99", "#B15928"],
201
+ "Pastel1": ["#FBB4AE", "#B3CDE3", "#CCEBC5", "#DECBE4", "#FED9A6", "#FFFFCC", "#E5D8BD", "#FDDAEC", "#F2F2F2"],
202
+ "Pastel2": ["#B3E2CD", "#FDCDAC", "#CBD5E8", "#F4CAE4", "#E6F5C9", "#FFF2AE", "#F1E2CC", "#CCCCCC"],
203
+ "Accent": ["#7FC97F", "#BEAED4", "#FDC086", "#FFFF99", "#386CB0", "#F0027F", "#BF5B17", "#666666"],
204
+ "Blues": ["#DEEBF7", "#C6DBEF", "#9ECAE1", "#6BAED6", "#4292C6", "#2171B5", "#08519C", "#08306B"],
205
+ "Greens": ["#E5F5E0", "#C7E9C0", "#A1D99B", "#74C476", "#41AB5D", "#238B45", "#006D2C", "#00441B"],
206
+ "Reds": ["#FEE0D2", "#FCBBA1", "#FC9272", "#FB6A4A", "#EF3B2C", "#CB181D", "#A50F15", "#67000D"],
207
+ "Oranges": ["#FEE6CE", "#FDD0A2", "#FDAE6B", "#FD8D3C", "#F16913", "#D94801", "#A63603", "#7F2704"],
208
+ "Purples": ["#EFEDF5", "#DADAEB", "#BCBDDC", "#9E9AC8", "#807DBA", "#6A51A3", "#54278F", "#3F007D"],
209
+ "Greys": ["#F0F0F0", "#D9D9D9", "#BDBDBD", "#969696", "#737373", "#525252", "#252525", "#000000"],
210
+ "RdBu": ["#B2182B", "#D6604D", "#F4A582", "#FDDBC7", "#D1E5F0", "#92C5DE", "#4393C3", "#2166AC"],
211
+ "PuOr": ["#B35806", "#E08214", "#FDB863", "#FEE0B6", "#D8DAEB", "#B2ABD2", "#8073AC", "#542788"],
212
+ }
213
+ _SEQUENTIAL = {"Blues", "Greens", "Reds", "Oranges", "Purples", "Greys", "RdBu", "PuOr"}
214
+ # Okabe-Ito: distinguishable with every common colour-vision deficiency.
215
+ OKABE_ITO = ["#E69F00", "#56B4E9", "#009E73", "#F0E442", "#0072B2", "#D55E00", "#CC79A7", "#000000"]
216
+
217
+
218
+ def to_hex(colour) -> str:
219
+ """Any colour plot3 accepts ("grey50", "steelblue", "#abc") as #rrggbb."""
220
+ from plot3.static import _named_colour
221
+
222
+ if not isinstance(colour, str):
223
+ return colour
224
+ text = colour.strip()
225
+ if text.startswith("#"):
226
+ body = text[1:]
227
+ return "#" + ("".join(ch * 2 for ch in body) if len(body) == 3 else body).lower()
228
+ named = _named_colour(text)
229
+ if named is None:
230
+ raise ValueError(
231
+ f"colour {colour!r} is not a CSS name, an R grey such as 'grey50', or #rrggbb"
232
+ )
233
+ return f"#{named}"
234
+
235
+
236
+ def _hex_rgb(colour: str) -> tuple[float, float, float]:
237
+ from plot3.static import _rgb
238
+
239
+ r, g, b = _rgb(colour)
240
+ return r / 255.0, g / 255.0, b / 255.0
241
+
242
+
243
+ def _rgb_hex(rgb) -> str:
244
+ return "#" + "".join(f"{int(round(max(0.0, min(1.0, c)) * 255)):02x}" for c in rgb)
245
+
246
+
247
+ def ramp_at(stops: list[str], t: float) -> str:
248
+ """Colour at ``t`` (0..1) along evenly spaced ``stops``."""
249
+ t = max(0.0, min(1.0, float(t)))
250
+ if len(stops) == 1:
251
+ return stops[0]
252
+ pos = t * (len(stops) - 1)
253
+ i = min(int(pos), len(stops) - 2)
254
+ frac = pos - i
255
+ a, b = _hex_rgb(stops[i]), _hex_rgb(stops[i + 1])
256
+ return _rgb_hex([a[k] + (b[k] - a[k]) * frac for k in range(3)])
257
+
258
+
259
+ def _sample(stops: list[str], n: int, begin: float = 0.0, end: float = 1.0) -> list[str]:
260
+ if n <= 1:
261
+ return [ramp_at(stops, (begin + end) / 2)]
262
+ return [ramp_at(stops, begin + (end - begin) * i / (n - 1)) for i in range(n)]
263
+
264
+
265
+ def extend_palette(base: list[str], n: int) -> list[str]:
266
+ """``n`` colours: the theme's own, then evenly spaced hues (ggplot2's hue
267
+ wheel) instead of an error past eight groups."""
268
+ if n <= len(base):
269
+ return list(base[:n])
270
+ extra = n - len(base)
271
+ hues = [(0.07 + i / extra) % 1.0 for i in range(extra)]
272
+ return list(base) + [_rgb_hex(colorsys.hls_to_rgb(h, 0.55, 0.62)) for h in hues]
273
+
274
+
275
+ class ColourScale:
276
+ """A discrete palette or a continuous ramp for the colour/fill channel."""
277
+
278
+ def __init__(self, kind, *, name=None, values=None, palette_fn=None,
279
+ breaks=None, labels=None, na_value="#7f7f7f",
280
+ low=None, mid=None, high=None, midpoint=None, limits=None,
281
+ stops=None, positions=None, identity=False):
282
+ self.kind = kind # "discrete" | "continuous"
283
+ self.stops = None if stops is None else [to_hex(c) for c in stops]
284
+ self.positions = None if positions is None else [float(v) for v in positions]
285
+ # scale_*_identity(): the column holds the colours themselves.
286
+ self.identity = bool(identity)
287
+ self.name = name
288
+ self.values = values
289
+ self.palette_fn = palette_fn
290
+ self.breaks = None if breaks is None else [str(b) for b in breaks]
291
+ self.labels = labels
292
+ self.na_value = na_value
293
+ self.low, self.mid, self.high, self.midpoint = low, mid, high, midpoint
294
+ self.limits = limits
295
+
296
+ def colours(self, levels: list[str], default: list[str]) -> list[str]:
297
+ """One colour per level, in level order, as #rrggbb."""
298
+ return [to_hex(c) for c in self._colours(levels, default)]
299
+
300
+ def _colours(self, levels: list[str], default: list[str]) -> list[str]:
301
+ if self.identity:
302
+ return [self.na_value if level in {"nan", "None", "<NA>"} else level for level in levels]
303
+ if isinstance(self.values, dict):
304
+ given = {str(k): v for k, v in self.values.items()}
305
+ return [given.get(level, self.na_value) for level in levels]
306
+ if self.values is not None:
307
+ values = list(self.values)
308
+ order = self.breaks or levels
309
+ by_level = {level: values[i] for i, level in enumerate(order) if i < len(values)}
310
+ if len(values) < len(levels):
311
+ raise ValueError(
312
+ f"scale_*_manual() has {len(values)} colours for {len(levels)} groups"
313
+ )
314
+ return [by_level.get(level, values[levels.index(level)]) for level in levels]
315
+ if self.palette_fn is not None:
316
+ return list(self.palette_fn(len(levels)))
317
+ return extend_palette(default, len(levels))
318
+
319
+ def ramp(self, lo: float, hi: float) -> list[str] | None:
320
+ """Continuous stops; a diverging scale centres ``mid`` on ``midpoint``."""
321
+ if self.kind != "continuous":
322
+ return None
323
+ if self.stops is not None:
324
+ if self.positions is None:
325
+ return list(self.stops)
326
+ # gradientn(values=): stops at those places along the scale.
327
+ rgb = np.array([_hex_rgb(c) for c in self.stops])
328
+ where = np.asarray(self.positions, dtype=np.float64)
329
+ grid = np.linspace(0.0, 1.0, 65)
330
+ return [_rgb_hex([np.interp(u, where, rgb[:, k]) for k in range(3)]) for u in grid]
331
+ if self.mid is None:
332
+ return [to_hex(self.low), to_hex(self.high)]
333
+ midpoint = 0.0 if self.midpoint is None else float(self.midpoint)
334
+ span = max(hi - lo, 1e-12)
335
+ centre = min(max((midpoint - lo) / span, 0.0), 1.0)
336
+ out = []
337
+ for i in range(65):
338
+ t = i / 64
339
+ if t <= centre:
340
+ u = 0.5 * (t / centre) if centre > 0 else 0.5
341
+ else:
342
+ u = 0.5 + 0.5 * ((t - centre) / (1 - centre)) if centre < 1 else 1.0
343
+ out.append(ramp_at([self.low, self.mid, self.high], u))
344
+ return out
345
+
346
+ def legend_entries(self, levels: list[str], colours: list[str]) -> list[dict]:
347
+ order = self.breaks if self.breaks is not None else levels
348
+ by_level = dict(zip(levels, colours))
349
+ out = []
350
+ for i, level in enumerate(order):
351
+ if level not in by_level:
352
+ continue
353
+ label = level
354
+ if isinstance(self.labels, dict):
355
+ label = str(self.labels.get(level, level))
356
+ elif isinstance(self.labels, (list, tuple)) and i < len(self.labels):
357
+ label = str(self.labels[i])
358
+ elif callable(self.labels):
359
+ label = str(self.labels(level))
360
+ out.append({"label": label, "color": by_level[level], "_level": level})
361
+ return out
362
+
363
+
364
+ def _discrete(name, values=None, palette_fn=None, breaks=None, labels=None, na_value="#7f7f7f"):
365
+ return ColourScale("discrete", name=name, values=values, palette_fn=palette_fn,
366
+ breaks=breaks, labels=labels, na_value=na_value)
367
+
368
+
369
+ def scale_colour_manual(values, *, breaks=None, labels=None, name=None, na_value="#7f7f7f"):
370
+ """Your colours: a list in level order, or ``{"level": "colour"}``."""
371
+ return _discrete(name, values=values, breaks=breaks, labels=labels, na_value=na_value)
372
+
373
+
374
+ def scale_colour_brewer(palette="Set1", *, direction=1, breaks=None, labels=None, name=None):
375
+ """ColorBrewer palettes: Set1, Set2, Set3, Dark2, Paired, Pastel1, Pastel2,
376
+ Accent (qualitative); Blues, Greens, Reds, Oranges, Purples, Greys, RdBu,
377
+ PuOr (ordered)."""
378
+ if palette not in _BREWER:
379
+ raise ValueError(f"palette {palette!r} is not one of {sorted(_BREWER)}")
380
+ stops = _BREWER[palette]
381
+
382
+ def fn(n):
383
+ if palette in _SEQUENTIAL:
384
+ out = _sample(stops, n, 0.15 if n < len(stops) else 0.0, 1.0)
385
+ else:
386
+ out = extend_palette(stops, n)
387
+ return out[::-1] if direction == -1 else out
388
+
389
+ return _discrete(name, palette_fn=fn, breaks=breaks, labels=labels)
390
+
391
+
392
+ def scale_colour_okabe_ito(*, breaks=None, labels=None, name=None):
393
+ """Okabe-Ito: eight colours safe for colour-blind readers."""
394
+ return _discrete(name, palette_fn=lambda n: extend_palette(OKABE_ITO, n), breaks=breaks, labels=labels)
395
+
396
+
397
+ def scale_colour_viridis_d(option="viridis", *, begin=0.0, end=1.0, direction=1,
398
+ breaks=None, labels=None, name=None):
399
+ """Evenly spaced colours from viridis (or magma, turbo) for categories."""
400
+ from plot3.themes import _CONT_PALETTES
401
+
402
+ if option not in _CONT_PALETTES:
403
+ raise ValueError(f"option {option!r} is not one of {sorted(_CONT_PALETTES)}")
404
+ stops = _CONT_PALETTES[option]
405
+
406
+ def fn(n):
407
+ out = _sample(stops, n, begin, end)
408
+ return out[::-1] if direction == -1 else out
409
+
410
+ return _discrete(name, palette_fn=fn, breaks=breaks, labels=labels)
411
+
412
+
413
+ def scale_colour_grey(start=0.2, end=0.8, *, breaks=None, labels=None, name=None):
414
+ """Greys for black-and-white print (0 black, 1 white)."""
415
+ def fn(n):
416
+ return [_rgb_hex([start + (end - start) * (i / max(n - 1, 1))] * 3) for i in range(n)]
417
+
418
+ return _discrete(name, palette_fn=fn, breaks=breaks, labels=labels)
419
+
420
+
421
+ def scale_colour_gradient(low="#132B43", high="#56B1F7", *, limits=None, name=None):
422
+ """A continuous ramp from ``low`` to ``high`` (ggplot2's default blues)."""
423
+ return ColourScale("continuous", name=name, low=low, high=high, limits=limits)
424
+
425
+
426
+ def scale_colour_gradient2(low="#832424", mid="#FFFFFF", high="#3A3A98", *, midpoint=0.0,
427
+ limits=None, name=None):
428
+ """Diverging: ``low`` below ``midpoint``, ``mid`` at it, ``high`` above."""
429
+ return ColourScale("continuous", name=name, low=low, mid=mid, high=high,
430
+ midpoint=midpoint, limits=limits)
431
+
432
+
433
+ def scale_colour_gradientn(colours=None, *, values=None, limits=None, name=None, colors=None):
434
+ """A continuous ramp through several ``colours``; ``values`` (0..1, one
435
+ per colour) places them along the scale."""
436
+ stops = colours if colours is not None else colors
437
+ if not stops or len(stops) < 2:
438
+ raise ValueError("scale_colour_gradientn() needs at least two colours")
439
+ if values is not None and len(values) != len(stops):
440
+ raise ValueError("scale_colour_gradientn(values=) needs one value per colour")
441
+ return ColourScale("continuous", name=name, stops=list(stops), positions=values, limits=limits)
442
+
443
+
444
+ def scale_colour_distiller(palette="Blues", *, direction=-1, limits=None, name=None):
445
+ """A ColorBrewer palette stretched over a number (ggplot2's distiller).
446
+ ``direction=-1`` (the default, as in ggplot2) puts the darkest colour at
447
+ the low end."""
448
+ if palette not in _BREWER:
449
+ raise ValueError(f"palette {palette!r} is not one of {sorted(_BREWER)}")
450
+ stops = list(_BREWER[palette])
451
+ if direction == -1:
452
+ stops = stops[::-1]
453
+ return ColourScale("continuous", name=name, stops=stops, limits=limits)
454
+
455
+
456
+ def scale_colour_identity(*, name=None, na_value="#7f7f7f"):
457
+ """Use the column's own colours ("red", "#1b9e77"), with no legend."""
458
+ return ColourScale("discrete", name=name, identity=True, na_value=na_value)
459
+
460
+
461
+ def hcl_hex(h: float, c: float, l: float) -> str:
462
+ """R's hcl(): polar CIE-LUV (D65) to sRGB, out-of-gamut values clipped."""
463
+
464
+ xn, yn, zn = 95.047, 100.0, 108.883
465
+ un = 4 * xn / (xn + 15 * yn + 3 * zn)
466
+ vn = 9 * yn / (xn + 15 * yn + 3 * zn)
467
+ if l <= 0:
468
+ return "#000000"
469
+ u = c * math.cos(math.radians(h))
470
+ v = c * math.sin(math.radians(h))
471
+ y = yn * (((l + 16) / 116) ** 3 if l > 8 else l / (24389 / 27))
472
+ up, vp = u / (13 * l) + un, v / (13 * l) + vn
473
+ x = 9.0 * y * up / (4 * vp)
474
+ z = -x / 3 - 5 * y + 3 * y / vp
475
+ x, y, z = x / 100, y / 100, z / 100
476
+ lin = (3.240479 * x - 1.537150 * y - 0.498535 * z,
477
+ -0.969256 * x + 1.875992 * y + 0.041556 * z,
478
+ 0.055648 * x - 0.204043 * y + 1.057311 * z)
479
+
480
+ def gamma(ch):
481
+ ch = min(max(ch, 0.0), 1.0)
482
+ return 12.92 * ch if ch <= 0.0031308 else 1.055 * ch ** (1 / 2.4) - 0.055
483
+
484
+ return "#" + "".join(f"{int(round(gamma(ch) * 255)):02x}" for ch in lin)
485
+
486
+
487
+ def scale_colour_hue(*, h=(15, 375), c=100, l=65, h_start=0, direction=1,
488
+ breaks=None, labels=None, name=None):
489
+ """ggplot2's default discrete colours: evenly spaced hues at one
490
+ chroma and lightness (#F8766D, #00BA38, #619CFF for three groups)."""
491
+ lo, hi = float(h[0]), float(h[1])
492
+
493
+ def fn(n):
494
+ top = hi - 360.0 / n if (hi - lo) % 360 < 1 else hi
495
+ hues = [lo + (top - lo) * i / max(n - 1, 1) for i in range(n)] if n > 1 else [lo]
496
+ hues = [(value + h_start) % 360 for value in hues]
497
+ out = [hcl_hex(value, c, l) for value in hues]
498
+ return out[::-1] if direction == -1 else out
499
+
500
+ return _discrete(name, palette_fn=fn, breaks=breaks, labels=labels)
501
+
502
+
503
+ scale_fill_hue = scale_colour_hue
504
+ scale_color_hue = scale_colour_hue
505
+ scale_fill_gradientn = scale_colour_gradientn
506
+ scale_color_gradientn = scale_colour_gradientn
507
+ scale_fill_distiller = scale_colour_distiller
508
+ scale_color_distiller = scale_colour_distiller
509
+ scale_fill_identity = scale_colour_identity
510
+ scale_color_identity = scale_colour_identity
511
+ scale_fill_manual = scale_colour_manual
512
+ scale_fill_brewer = scale_colour_brewer
513
+ scale_fill_okabe_ito = scale_colour_okabe_ito
514
+ scale_fill_viridis_d = scale_colour_viridis_d
515
+ scale_fill_grey = scale_colour_grey
516
+ scale_fill_gradient = scale_colour_gradient
517
+ scale_fill_gradient2 = scale_colour_gradient2
518
+ scale_color_manual = scale_colour_manual
519
+ scale_color_brewer = scale_colour_brewer
520
+ scale_color_okabe_ito = scale_colour_okabe_ito
521
+ scale_color_viridis_d = scale_colour_viridis_d
522
+ scale_color_grey = scale_colour_grey
523
+ scale_color_gradient = scale_colour_gradient
524
+ scale_color_gradient2 = scale_colour_gradient2
525
+
526
+
527
+ # ── shape and linetype ───────────────────────────────────────────────────────
528
+
529
+
530
+ class KeyScale:
531
+ """Manual values for aes(shape=) or aes(linetype=)."""
532
+
533
+ def __init__(self, aesthetic, values, breaks=None, labels=None, name=None):
534
+ self.aesthetic = aesthetic
535
+ self.values = list(values)
536
+ self.breaks = None if breaks is None else [str(b) for b in breaks]
537
+ self.labels = labels
538
+ self.name = name
539
+ if aesthetic == "shape":
540
+ from plot3.geoms import shape_name
541
+
542
+ self.values = [shape_name(v) for v in self.values]
543
+ else:
544
+ from plot3.geoms import dash_pattern
545
+
546
+ for v in self.values:
547
+ dash_pattern(v)
548
+
549
+ def value_for(self, index: int, level: str):
550
+ order = self.breaks or []
551
+ if level in order and order.index(level) < len(self.values):
552
+ return self.values[order.index(level)]
553
+ return self.values[index % len(self.values)]
554
+
555
+
556
+ def scale_shape_manual(values, *, breaks=None, labels=None, name=None):
557
+ """Symbols per level: names ("triangle") or R numbers (17)."""
558
+ return KeyScale("shape", values, breaks, labels, name)
559
+
560
+
561
+ def scale_linetype_manual(values, *, breaks=None, labels=None, name=None):
562
+ """Dash patterns per level: "solid", "dashed", "dotted", … or hex "44"."""
563
+ return KeyScale("linetype", values, breaks, labels, name)
564
+
565
+
566
+ # ── size and alpha ───────────────────────────────────────────────────────────
567
+
568
+
569
+ class SizeScale:
570
+ """How a number maps to point size: by area across a ``range``
571
+ (scale_size), or by area from zero (scale_size_area)."""
572
+
573
+ def __init__(self, kind, *, range=None, max_size=None, limits=None, breaks=None, name=None):
574
+ self.kind = kind # "range" | "area"
575
+ self.range = None if range is None else (float(range[0]), float(range[1]))
576
+ self.max_size = None if max_size is None else float(max_size)
577
+ self.limits = None if limits is None else (float(limits[0]), float(limits[1]))
578
+ self.breaks = None if breaks is None else [float(b) for b in breaks]
579
+ self.name = name
580
+ if self.range is not None and not (0 <= self.range[0] <= self.range[1] and self.range[1] > 0):
581
+ raise ValueError("scale_size(range=) is (smallest, largest), for example (4, 23)")
582
+
583
+
584
+ def scale_size(name=None, *, range=(4.0, 23.0), limits=None, breaks=None):
585
+ """Point area across ``range``: the smallest value gets the first size,
586
+ the largest the second (ggplot2's default size scale). Sizes are in the
587
+ units of ``geom_point(size=)`` (pixels in 2D); (4, 23) is ggplot2's
588
+ ``range = c(1, 6)``."""
589
+ return SizeScale("range", range=range, limits=limits, breaks=breaks, name=name)
590
+
591
+
592
+ def scale_size_area(name=None, *, max_size=23.0, breaks=None):
593
+ """Point area in proportion to the value, zero at zero (plot3's default,
594
+ with ``max_size`` the largest point)."""
595
+ return SizeScale("area", max_size=max_size, breaks=breaks, name=name)
596
+
597
+
598
+ class AlphaScale:
599
+ """How a number maps to opacity: across ``range`` (0..1)."""
600
+
601
+ def __init__(self, *, range=(0.1, 1.0), limits=None, name=None):
602
+ lo, hi = float(range[0]), float(range[1])
603
+ if not (0.0 <= lo <= 1.0 and 0.0 <= hi <= 1.0):
604
+ raise ValueError("scale_alpha(range=) values are opacities from 0 to 1")
605
+ self.range = (lo, hi)
606
+ self.limits = None if limits is None else (float(limits[0]), float(limits[1]))
607
+ self.name = name
608
+
609
+
610
+ def scale_alpha(name=None, *, range=(0.1, 1.0), limits=None):
611
+ """Opacity for ``aes(alpha=)``: the smallest value is ``range[0]``, the
612
+ largest ``range[1]``, as in ggplot2."""
613
+ return AlphaScale(range=range, limits=limits, name=name)
614
+
615
+
616
+ scale_alpha_continuous = scale_alpha
617
+
618
+
619
+ __all__ = [
620
+ "scale_x_continuous", "scale_y_continuous", "scale_x_reverse", "scale_y_reverse",
621
+ "scale_x_discrete", "scale_y_discrete", "scale_x_date", "scale_y_date",
622
+ "scale_x_datetime", "scale_y_datetime", "xlim", "ylim", "lims",
623
+ "scale_colour_manual", "scale_fill_manual", "scale_color_manual",
624
+ "scale_colour_brewer", "scale_fill_brewer", "scale_color_brewer",
625
+ "scale_colour_okabe_ito", "scale_fill_okabe_ito", "scale_color_okabe_ito",
626
+ "scale_colour_viridis_d", "scale_fill_viridis_d", "scale_color_viridis_d",
627
+ "scale_colour_grey", "scale_fill_grey", "scale_color_grey",
628
+ "scale_colour_gradient", "scale_fill_gradient", "scale_color_gradient",
629
+ "scale_colour_gradient2", "scale_fill_gradient2", "scale_color_gradient2",
630
+ "scale_shape_manual", "scale_linetype_manual",
631
+ "scale_colour_gradientn", "scale_fill_gradientn", "scale_color_gradientn",
632
+ "scale_colour_distiller", "scale_fill_distiller", "scale_color_distiller",
633
+ "scale_colour_identity", "scale_fill_identity", "scale_color_identity",
634
+ "scale_size", "scale_size_area", "scale_alpha", "scale_alpha_continuous",
635
+ "scale_colour_hue", "scale_fill_hue", "scale_color_hue",
636
+ ]