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/scales.py ADDED
@@ -0,0 +1,387 @@
1
+ """Positional scales, ticks, and column typing."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ import numpy as np
8
+ import pandas as pd
9
+
10
+
11
+ def resolution(x, *, zero: bool = True) -> float:
12
+ """Smallest non-zero distance between adjacent values (ggplot2 ``resolution``).
13
+
14
+ Used to turn a relative bar ``width`` (default 0.9) into data units:
15
+ ``data_width = resolution(x) * width``.
16
+
17
+ * Integer-like vectors → ``1`` (ggplot2 treats integers as unit-spaced).
18
+ * Single unique value / zero range → ``1``.
19
+ * Otherwise → minimum positive difference between sorted unique values.
20
+ * If ``zero`` is True (ggplot2 default for some paths), ``0`` is included
21
+ in the unique set before measuring gaps.
22
+ """
23
+ raw = np.asarray(x)
24
+ # Integer storage (int32/Int64/…) matches ggplot2 is.integer → 1.
25
+ if np.issubdtype(raw.dtype, np.integer):
26
+ return 1.0
27
+ arr = np.asarray(raw, dtype=np.float64).ravel()
28
+ arr = arr[np.isfinite(arr)]
29
+ if arr.size == 0:
30
+ return 1.0
31
+ lo = float(np.min(arr))
32
+ hi = float(np.max(arr))
33
+ if hi <= lo or not math.isfinite(lo) or not math.isfinite(hi):
34
+ return 1.0
35
+ uniq = np.unique(arr)
36
+ if zero:
37
+ uniq = np.unique(np.concatenate([uniq, np.asarray([0.0])]))
38
+ if uniq.size < 2:
39
+ return 1.0
40
+ d = np.diff(np.sort(uniq))
41
+ tol = math.sqrt(np.finfo(float).eps)
42
+ positive = d[d > tol]
43
+ if positive.size == 0:
44
+ return 1.0
45
+ return float(np.min(positive))
46
+
47
+
48
+ def _nearest_multiple(value: float, step: float) -> float:
49
+ """The multiple of ``step`` closest to ``value``. A tie takes the higher one."""
50
+ k = value / step
51
+ down = math.floor(k + 1e-10)
52
+ up = math.ceil(k - 1e-10)
53
+ if abs(k - down) < abs(up - k) - 1e-10:
54
+ return down * step
55
+ return up * step
56
+
57
+
58
+ def _ticks_on_step(lo: float, hi: float, step: float) -> list[float] | None:
59
+ start = _nearest_multiple(lo, step)
60
+ end = _nearest_multiple(hi, step)
61
+ if end < start:
62
+ start, end = end, start
63
+ count = int(round((end - start) / step))
64
+ if count < 0 or count > 40:
65
+ return None
66
+ origin = int(round(start / step))
67
+ out = []
68
+ for i in range(count + 1):
69
+ t = (origin + i) * step
70
+ if abs(t) < abs(step) * 1e-8:
71
+ t = 0.0
72
+ else:
73
+ t = float(f"{t:.12g}")
74
+ out.append(t)
75
+ return out
76
+
77
+
78
+ def _classic_ticks(lo: float, hi: float, n: int) -> list[float]:
79
+ raw = (hi - lo) / max(1, n)
80
+ mag = 10 ** math.floor(math.log10(max(raw, 1e-12)))
81
+ step = 10 * mag
82
+ for mult in (1, 2, 5, 10):
83
+ if raw <= mult * mag:
84
+ step = mult * mag
85
+ break
86
+ found = _ticks_on_step(lo, hi, step)
87
+ return found or [lo, hi]
88
+
89
+
90
+ def nice_ticks(lo: float, hi: float, n: int = 5) -> list[float]:
91
+ """About ``n`` ticks on a 1-2-2.5-5 grid inside the range, as ggplot2's
92
+ extended breaks: 0, 2, 4, 6, 8 for 0 to 9.2, not every unit.
93
+
94
+ Ticks may sit up to 5% of the span past either end (the view's margin),
95
+ so a sample maximum of 0.9997 still gets a tick at 1. Among steps with
96
+ about ``n`` ticks, the one whose ticks reach nearest both ends wins;
97
+ 2.5 steps count as a little less round than 1, 2, and 5.
98
+ """
99
+ if not math.isfinite(lo) or not math.isfinite(hi) or hi <= lo:
100
+ return [lo]
101
+ span = hi - lo
102
+ target = max(3, int(n))
103
+ pad = 0.05 * span
104
+ raw = span / max(1, target - 1)
105
+ exp = math.floor(math.log10(max(raw, 1e-300)))
106
+ best: list[float] | None = None
107
+ best_score = math.inf
108
+ best_step = -1.0
109
+ for shift in (-1, 0, 1):
110
+ base = 10.0 ** (exp + shift)
111
+ for mult in (1, 2, 2.5, 5):
112
+ step = mult * base
113
+ first = math.ceil((lo - pad) / step - 1e-9)
114
+ last = math.floor((hi + pad) / step + 1e-9)
115
+ count = last - first + 1
116
+ if count < 2 or count > 40:
117
+ continue
118
+ score = float(abs(count - target))
119
+ if count < 4:
120
+ score += 3.0
121
+ if count > target + 2:
122
+ score += 2.0 * (count - target - 2)
123
+ gap = max(0.0, first * step - lo) + max(0.0, hi - last * step)
124
+ score += 2.0 * gap / span
125
+ if mult == 2.5:
126
+ score += 0.25
127
+ if score < best_score - 1e-9 or (abs(score - best_score) <= 1e-9 and step > best_step):
128
+ # Round to the step, not to significant figures, so ticks
129
+ # near 1e12 stay apart.
130
+ digits = max(0, 3 - math.floor(math.log10(step)))
131
+ best = []
132
+ for k in range(first, last + 1):
133
+ tick = k * step
134
+ best.append(0.0 if abs(tick) < abs(step) * 1e-8 else round(tick, digits))
135
+ best_score = score
136
+ best_step = step
137
+ if best:
138
+ return best
139
+ return _classic_ticks(lo, hi, n)
140
+
141
+
142
+ def fmt_num(v: float) -> str:
143
+ if v == 0:
144
+ return "0"
145
+ a = abs(v)
146
+ if a >= 1e6 or a < 1e-4:
147
+ return f"{v:.3g}"
148
+ s = f"{v:.6f}".rstrip("0").rstrip(".")
149
+ return s
150
+
151
+
152
+ def fmt_ticks(values) -> list[str]:
153
+ """Tick labels that stay apart: 1000000000001, not five "1e+12"."""
154
+ labels = [fmt_num(v) for v in values]
155
+ if len(set(labels)) == len(labels):
156
+ return labels
157
+ if all(float(v).is_integer() and abs(v) < 1e15 for v in values):
158
+ return [str(int(v)) for v in values]
159
+ for digits in range(4, 16):
160
+ labels = [f"{v:.{digits}g}" for v in values]
161
+ if len(set(labels)) == len(labels):
162
+ return labels
163
+ return labels
164
+
165
+
166
+ def log_ticks(lo: float, hi: float) -> list[list]:
167
+ """Ticks for a log10 scale.
168
+
169
+ ``lo`` and ``hi`` are already log10 values (the scale's stored domain).
170
+ Each tick is ``[log10(value), label]`` so it sits in that same space.
171
+ """
172
+ if not math.isfinite(lo) or not math.isfinite(hi) or hi <= lo:
173
+ val = 10 ** lo if math.isfinite(lo) else 1.0
174
+ return [[lo, fmt_num(val)]]
175
+ span = hi - lo
176
+ mults = (1, 2, 5) if span <= 3 else (1,)
177
+ exp0 = math.floor(lo)
178
+ exp1 = math.ceil(hi)
179
+ out: list[list] = []
180
+ for exp in range(int(exp0), int(exp1) + 1):
181
+ for mult in mults:
182
+ val = mult * 10.0 ** exp
183
+ lv = math.log10(val)
184
+ if lv < lo - 1e-9 or lv > hi + 1e-9:
185
+ continue
186
+ out.append([lv, fmt_num(val)])
187
+ if not out:
188
+ out = [[lo, fmt_num(10 ** lo)], [hi, fmt_num(10 ** hi)]]
189
+ return out
190
+
191
+
192
+ DT_LADDERS = [
193
+ # (span_seconds >, [(pandas freq, strftime fmt), coarse -> fine])
194
+ (2 * 365 * 86400, [("YS", "%Y"), ("QS", "%b %Y"), ("MS", "%b %Y")]),
195
+ (90 * 86400, [("MS", "%b %Y"), ("W", "%b %d"), ("D", "%b %d")]),
196
+ (3 * 86400, [("D", "%b %d"), ("6h", "%d %Hh"), ("h", "%H:%M")]),
197
+ (3 * 3600, [("h", "%H:%M"), ("15min", "%H:%M"), ("min", "%H:%M")]),
198
+ (0, [("min", "%H:%M"), ("15s", "%H:%M:%S"), ("s", "%H:%M:%S")]),
199
+ ]
200
+
201
+
202
+ def dt_ladder(lo_s: float, hi_s: float) -> list[list[list]]:
203
+ """3-level [position_seconds, label] ladders; JS picks by visible count."""
204
+ span = hi_s - lo_s
205
+ for min_span, freqs in DT_LADDERS:
206
+ if span > min_span:
207
+ break
208
+ lo_ts = pd.Timestamp(lo_s, unit="s")
209
+ hi_ts = pd.Timestamp(hi_s, unit="s")
210
+ ladder = []
211
+ for freq, fmt in freqs:
212
+ try:
213
+ idx = pd.date_range(lo_ts.floor("s"), hi_ts.ceil("s"), freq=freq)
214
+ except Exception:
215
+ idx = pd.DatetimeIndex([lo_ts, hi_ts])
216
+ if len(idx) > 400:
217
+ idx = idx[:: len(idx) // 400 + 1]
218
+ ladder.append(
219
+ [[t.timestamp(), t.strftime(fmt)] for t in idx]
220
+ )
221
+ return ladder
222
+
223
+
224
+ class Scale:
225
+ """Resolved positional scale: numeric, datetime or categorical.
226
+
227
+ ``trans`` is ``None`` or ``"log10"``. A log scale stores ``lo`` / ``hi``
228
+ and tick positions in log10 space. Labels stay in the original units.
229
+ """
230
+
231
+ def __init__(self, kind: str, trans: str | None = None):
232
+ self.kind = kind # "num" | "dt" | "cat"
233
+ self.trans = trans
234
+ self.lo = math.inf
235
+ self.hi = -math.inf
236
+ self.cats: list[str] = []
237
+
238
+ def widen(self, values: np.ndarray):
239
+ if len(values) == 0:
240
+ return
241
+ finite = np.asarray(values, dtype=np.float64)
242
+ finite = finite[np.isfinite(finite)]
243
+ if finite.size == 0:
244
+ return
245
+ self.lo = min(self.lo, float(finite.min()))
246
+ self.hi = max(self.hi, float(finite.max()))
247
+
248
+ def finish(self):
249
+ if self.kind == "cat":
250
+ self.lo, self.hi = -0.5, max(0.5, len(self.cats) - 0.5)
251
+ elif not math.isfinite(self.lo):
252
+ self.lo, self.hi = 0.0, 1.0
253
+ elif self.hi <= self.lo:
254
+ self.lo, self.hi = self.lo - 0.5, self.hi + 0.5
255
+
256
+ def spec(self) -> dict:
257
+ d = {"kind": self.kind, "lo": self.lo, "hi": self.hi}
258
+ lo, hi = min(self.lo, self.hi), max(self.lo, self.hi) # a reversed axis
259
+ if self.trans:
260
+ d["trans"] = self.trans
261
+ if self.kind == "cat":
262
+ d["cats"] = self.cats
263
+ elif self.kind == "dt":
264
+ d["ladder"] = dt_ladder(lo, hi)
265
+ elif self.trans == "log10":
266
+ d["ticks"] = log_ticks(lo, hi)
267
+ else:
268
+ ticks = nice_ticks(lo, hi)
269
+ d["ticks"] = [list(pair) for pair in zip(ticks, fmt_ticks(ticks))]
270
+ custom = getattr(self, "custom", None)
271
+ if custom is not None:
272
+ _apply_custom(d, self, custom, lo, hi)
273
+ return d
274
+
275
+
276
+ def _is_na(value) -> bool:
277
+ try:
278
+ return value is None or bool(pd.isna(value))
279
+ except (TypeError, ValueError):
280
+ return False
281
+
282
+
283
+ def ordered_levels(values, *, keep_order: bool = False) -> list:
284
+ """Distinct values in ggplot2's level order, missing values last.
285
+
286
+ ``keep_order`` (a categorical column) keeps the given order. Otherwise
287
+ booleans run False, True; numbers sort numerically (4, 6, 8, 10, not
288
+ "10" before "4"); anything else sorts as text.
289
+ """
290
+ values = list(values)
291
+ present = [v for v in values if not _is_na(v)]
292
+ missing = len(present) != len(values)
293
+ distinct = list(dict.fromkeys(present))
294
+ if not keep_order:
295
+ if all(isinstance(v, (bool, np.bool_)) for v in distinct):
296
+ distinct.sort(key=bool)
297
+ elif all(
298
+ isinstance(v, (int, float, np.integer, np.floating))
299
+ and not isinstance(v, (bool, np.bool_))
300
+ for v in distinct
301
+ ):
302
+ distinct.sort(key=float)
303
+ else:
304
+ distinct.sort(key=str)
305
+ return distinct + ([None] if missing else [])
306
+
307
+
308
+ def col_values(s: pd.Series) -> tuple[str, np.ndarray, list[str]]:
309
+ """Series -> (scale kind, float64 positions, categories)."""
310
+ if pd.api.types.is_bool_dtype(s):
311
+ # True/False are two groups (ggplot2), not a 0-1 colour gradient.
312
+ cats = [str(v) for v in ordered_levels(s.dropna().unique().tolist())]
313
+ idx = {c: i for i, c in enumerate(cats)}
314
+ return "cat", s.map(lambda v: idx.get(str(v), np.nan)).to_numpy(np.float64), cats
315
+ if pd.api.types.is_datetime64_any_dtype(s):
316
+ if getattr(s.dtype, "tz", None) is not None:
317
+ s = s.dt.tz_convert("UTC").dt.tz_localize(None)
318
+ # normalize the unit: pandas 3.0 defaults to us, not ns
319
+ v = s.astype("datetime64[ns]").astype("int64").to_numpy(np.float64)
320
+ return "dt", v / 1e9, []
321
+ if isinstance(s.dtype, pd.CategoricalDtype):
322
+ return "cat", s.cat.codes.to_numpy(np.float64), [str(c) for c in s.cat.categories]
323
+ if pd.api.types.is_numeric_dtype(s):
324
+ return "num", s.to_numpy(np.float64), []
325
+ cats = sorted(s.dropna().astype(str).unique().tolist()) # ggplot2 sorts
326
+ idx = {c: i for i, c in enumerate(cats)}
327
+ return "cat", s.astype(str).map(idx).to_numpy(np.float64), cats
328
+
329
+
330
+
331
+ def _apply_custom(d: dict, scale: "Scale", custom, lo: float, hi: float) -> None:
332
+ """scale_x_continuous(breaks=, labels=), scale_x_discrete(labels=), and
333
+ scale_x_date(date_breaks=, date_labels=) on a resolved scale spec.
334
+
335
+ ``fixed`` tells the viewer to keep these ticks when it zooms.
336
+ """
337
+ from plot3.scaling import date_freq, format_label
338
+
339
+ labels = custom.labels
340
+ if scale.kind == "cat":
341
+ if isinstance(labels, dict):
342
+ d["cats"] = [str(labels.get(c, c)) for c in scale.cats]
343
+ elif isinstance(labels, (list, tuple)):
344
+ d["cats"] = [str(labels[i]) if i < len(labels) else c for i, c in enumerate(scale.cats)]
345
+ elif callable(labels):
346
+ d["cats"] = [str(labels(c)) for c in scale.cats]
347
+ return
348
+ if scale.kind == "dt":
349
+ fmt = custom.date_labels
350
+ if custom.date_breaks:
351
+ start = pd.Timestamp(lo, unit="s")
352
+ stop = pd.Timestamp(hi, unit="s")
353
+ freq = date_freq(custom.date_breaks)
354
+ stamps = pd.date_range(start.floor("D"), stop, freq=freq)
355
+ stamps = [t for t in stamps if lo <= t.timestamp() <= hi]
356
+ d["ticks"] = [[t.timestamp(), t.strftime(fmt or "%Y-%m-%d")] for t in stamps]
357
+ d["fixed"] = True
358
+ elif fmt:
359
+ d["ladder"] = [
360
+ [[t, pd.Timestamp(t, unit="s").strftime(fmt)] for t, _label in level]
361
+ for level in d.get("ladder") or []
362
+ ]
363
+ return
364
+ log = scale.trans == "log10"
365
+ if custom.breaks is not None:
366
+ ticks = []
367
+ for i, value in enumerate(custom.breaks):
368
+ value = float(value)
369
+ if log and value <= 0:
370
+ continue
371
+ pos = math.log10(value) if log else value
372
+ if not lo - 1e-9 <= pos <= hi + 1e-9:
373
+ continue
374
+ if isinstance(labels, (list, tuple)):
375
+ text = str(labels[i])
376
+ elif labels is not None:
377
+ text = format_label(value, labels)
378
+ else:
379
+ text = fmt_num(value)
380
+ ticks.append([pos, text])
381
+ d["ticks"] = ticks
382
+ d["fixed"] = True
383
+ elif labels is not None and not isinstance(labels, (list, tuple)):
384
+ d["ticks"] = [
385
+ [pos, format_label(10 ** pos if log else pos, labels)] for pos, _text in d.get("ticks") or []
386
+ ]
387
+ d["fixed"] = True