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/function.py ADDED
@@ -0,0 +1,1301 @@
1
+ """Sample a :class:`plot3.expr.Formula` into a drawable line or surface.
2
+
3
+ ``geom_function`` stays a thin parameter holder. At build time this module
4
+ evaluates it and returns a ``geom_line``, ``geom_path``, or ``geom_surface``
5
+ with ``data_override`` already filled — the same path bar and density stats use.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from contextlib import contextmanager
11
+ from functools import partial
12
+ from typing import Any, Callable
13
+
14
+ import numpy as np
15
+ import pandas as pd
16
+
17
+ from plot3.contour import _contour_lines, _refine_active_cells
18
+ from plot3.expr import ExprError, Formula, _missing_param_build_message, evaluate
19
+ from plot3.mathtext import _number, split_math
20
+ from plot3.geoms import _Geom, aes, geom_line, geom_path
21
+ from plot3.special import support_hint
22
+ from plot3.stats3d import regular_grid_mesh
23
+ from plot3.table import has_column, numeric_array
24
+
25
+ _DEFAULT_DOMAIN = (-10.0, 10.0)
26
+ _N_CURVE = 501
27
+ _N_GRID = 80
28
+
29
+
30
+ def data_domains(figure: Any, data: Any) -> dict[str, tuple[float, float]]:
31
+ """Numeric ranges of x/y/z columns already mapped on non-function layers."""
32
+ if data is None:
33
+ return {}
34
+ found: dict[str, tuple[float, float]] = {}
35
+ mappings = [getattr(figure, "mapping", None) or {}]
36
+ for layer in getattr(figure, "layers", []):
37
+ if getattr(layer, "kind", None) == "function":
38
+ continue
39
+ mapping = getattr(layer, "mapping", None)
40
+ if mapping:
41
+ mappings.append(mapping)
42
+ for mapping in mappings:
43
+ getter = getattr(mapping, "get", None)
44
+ if getter is None:
45
+ continue
46
+ for axis in ("x", "y", "z"):
47
+ if axis in found:
48
+ continue
49
+ column = getter(axis)
50
+ if not column or not has_column(data, column):
51
+ continue
52
+ try:
53
+ values = numeric_array(data, column, dropna=True)
54
+ except Exception:
55
+ continue
56
+ if values.size == 0:
57
+ continue
58
+ lo = float(np.min(values))
59
+ hi = float(np.max(values))
60
+ if not np.isfinite(lo) or not np.isfinite(hi):
61
+ continue
62
+ if hi <= lo:
63
+ lo, hi = lo - 1.0, hi + 1.0
64
+ found[axis] = (lo, hi)
65
+ return found
66
+
67
+
68
+ def expand_function(
69
+ geom: _Geom,
70
+ base_mapping: Any,
71
+ data: Any,
72
+ domains: dict[str, tuple[float, float]] | None = None,
73
+ transition: Any = None,
74
+ slider: Any = None,
75
+ coord: Any = None,
76
+ addons: list | None = None,
77
+ ) -> list:
78
+ """Turn ``geom_function`` into one or more line, area, or point layers."""
79
+ del base_mapping, data # the formula carries its own samples
80
+ formula: Formula = geom.formula
81
+ domains = domains or {}
82
+ # A slider replaces a transition. build_spec rejects having both.
83
+ sweep = slider if slider is not None else transition
84
+ ranges = getattr(sweep, "ranges", None) or {}
85
+ formula = _bind_swept_symbols(geom, formula, ranges)
86
+ pending = tuple(getattr(formula, "pending", ()) or ())
87
+ if pending:
88
+ missing = [name for name in pending if name not in ranges]
89
+ if missing:
90
+ raise ExprError(_missing_param_build_message(missing[0]))
91
+ namespace = getattr(formula, "namespace", None) or {}
92
+ # A coefficient passed as a keyword (a=2) already sits in the namespace.
93
+ # Naming it on a slider or transition sweeps that value instead.
94
+ covered = any(name in pending or name in namespace for name in ranges)
95
+ # Imported here: calculus imports this module at load time.
96
+ from plot3.calculus import animation_blocked, attach_calculus, expand_special
97
+
98
+ extras = list(addons or ())
99
+ blocked = animation_blocked(formula, geom, coord, extras)
100
+ if ranges and blocked:
101
+ raise ValueError(blocked)
102
+ special = expand_special(geom, formula, domains, coord)
103
+ if special is not None:
104
+ if extras or tuple(getattr(geom, "marks", ()) or ()):
105
+ _reject_curve_extras(geom, extras)
106
+ _link_colors(special)
107
+ return special
108
+ axes = _assign_axes(formula)
109
+ if ranges and (pending or covered):
110
+ return [_expand_animated(geom, formula, axes, domains, sweep)]
111
+ if axes.kind == "surface":
112
+ primary = _expand_surface(geom, formula, axes, domains)
113
+ elif axes.kind == "implicit":
114
+ primary = _expand_implicit(geom, formula, axes, domains)
115
+ else:
116
+ primary = _expand_curve(geom, formula, axes, domains)
117
+ layers = attach_calculus(primary, geom, formula, domains, extras)
118
+ _link_colors(layers)
119
+ return layers
120
+
121
+
122
+ def _bind_swept_symbols(geom, formula: Formula, ranges: dict) -> Formula:
123
+ """Re-read symbols a slider or transition sweeps as coefficients.
124
+
125
+ ``sin(x - t)`` parses ``t`` as a second plot variable and
126
+ ``dnorm(x, mu, 1)`` does the same with ``mu``, which would turn a
127
+ travelling wave into a static surface. Naming the symbol on
128
+ ``transition_time(t=...)`` says it is a coefficient. Each frame sets
129
+ its value; the low end of the range only stands in while parsing.
130
+ """
131
+ swept = [name for name in ranges if name in (formula.variables or ())]
132
+ source = getattr(geom, "_source", None)
133
+ if not swept or source is None or formula.mode == "callable":
134
+ return formula
135
+ params = dict(getattr(geom, "params", None) or {})
136
+ for name in swept:
137
+ params[name] = float(ranges[name][0])
138
+ from plot3.expr import parse_formula
139
+
140
+ return parse_formula(source, params, defer_missing=True)
141
+
142
+
143
+ def _reject_curve_extras(geom, addons) -> None:
144
+ """Parametric, polar, and inequality layers stay a single curve."""
145
+ if tuple(getattr(geom, "marks", ()) or ()):
146
+ raise ExprError(
147
+ "mark='roots' and mark='extrema' are for a curve y = f(x)"
148
+ )
149
+ name = type(addons[0]).__name__
150
+ raise ExprError(
151
+ f"{name}() is for a curve y = f(x). "
152
+ f'For example geom_function("y = x^2") + {name}(...)'
153
+ )
154
+
155
+
156
+ def _link_colors(layers: list) -> None:
157
+ """Point a fill at the curve it belongs to, so each curve keeps its colour."""
158
+ primary = next(
159
+ (layer for layer in layers if getattr(layer, "_formula_primary", False)),
160
+ None,
161
+ )
162
+ if primary is None:
163
+ return
164
+ token = f"formula-{id(primary)}"
165
+ primary._color_key = token
166
+ for layer in layers:
167
+ if layer is primary:
168
+ continue
169
+ if getattr(layer, "_inherit_color", False) and not getattr(layer, "_inherit_from", None):
170
+ layer._inherit_from = token
171
+
172
+
173
+ class _Axes:
174
+ def __init__(
175
+ self,
176
+ kind: str,
177
+ x: str,
178
+ y: str,
179
+ z: str | None,
180
+ computed: str | None,
181
+ ):
182
+ self.kind = kind
183
+ self.x = x
184
+ self.y = y
185
+ self.z = z
186
+ self.computed = computed
187
+
188
+
189
+ def _assign_axes(formula: Formula) -> _Axes:
190
+ variables = list(formula.variables)
191
+ if formula.mode == "callable":
192
+ if len(formula.fn_args) >= 2:
193
+ return _Axes("surface", formula.fn_args[0], formula.fn_args[1], "z", "z")
194
+ name = formula.fn_args[0] if formula.fn_args else "x"
195
+ y_name = "f" if name == "y" else "y"
196
+ return _Axes("curve", name, y_name, None, "y")
197
+ if formula.mode == "implicit":
198
+ ordered = _prefer_xy(variables)
199
+ return _Axes("implicit", ordered[0], ordered[1], None, None)
200
+ dependent = formula.dependent or "y"
201
+ if len(variables) >= 2:
202
+ ordered = _prefer_xy(variables)
203
+ z_name = dependent if dependent not in ordered else "z"
204
+ return _Axes("surface", ordered[0], ordered[1], z_name, "z")
205
+ if len(variables) == 1 and dependent == "x":
206
+ return _Axes("curve", "x", variables[0], None, "x")
207
+ if len(variables) == 0:
208
+ return _Axes("curve", "x", dependent, None, "y")
209
+ return _Axes("curve", variables[0], dependent, None, "y")
210
+
211
+
212
+ def _prefer_xy(names: list[str]) -> list[str]:
213
+ """Put ``x`` then ``y`` first when those names are present."""
214
+ rest = [name for name in names if name not in {"x", "y"}]
215
+ ordered: list[str] = []
216
+ for prefer in ("x", "y"):
217
+ if prefer in names:
218
+ ordered.append(prefer)
219
+ ordered.extend(rest)
220
+ return ordered
221
+
222
+
223
+ def _limit_pair(value: Any, name: str) -> tuple[float, float] | None:
224
+ if value is None:
225
+ return None
226
+ if (
227
+ isinstance(value, (tuple, list))
228
+ and len(value) == 2
229
+ and _both_real(value[0], value[1])
230
+ ):
231
+ lo, hi = float(value[0]), float(value[1])
232
+ if hi < lo:
233
+ lo, hi = hi, lo
234
+ if hi == lo:
235
+ hi = lo + 1.0
236
+ return lo, hi
237
+ raise ExprError(f"{name} must be a pair of numbers, for example (-2, 2)")
238
+
239
+
240
+ def _both_real(a: Any, b: Any) -> bool:
241
+ try:
242
+ return np.isfinite(float(a)) and np.isfinite(float(b))
243
+ except (TypeError, ValueError):
244
+ return False
245
+
246
+
247
+ def _sample_count(geom: _Geom, grid: bool) -> int:
248
+ chosen = getattr(geom, "n", None)
249
+ if chosen is None:
250
+ return _N_GRID if grid else _N_CURVE
251
+ count = int(chosen)
252
+ if count < 2:
253
+ raise ExprError("n must be at least 2")
254
+ return count
255
+
256
+
257
+ def _domain_for(
258
+ geom: _Geom,
259
+ axis: str,
260
+ domains: dict[str, tuple[float, float]],
261
+ ) -> tuple[tuple[float, float], str]:
262
+ """Return ``((lo, hi), source)`` where source is user, data, or default."""
263
+ explicit = _limit_pair(getattr(geom, axis + "lim", None), axis + "lim")
264
+ if explicit is not None:
265
+ return explicit, "user"
266
+ if axis in domains:
267
+ return domains[axis], "data"
268
+ return _DEFAULT_DOMAIN, "default"
269
+
270
+
271
+ def _density_domain(formula: Formula) -> tuple[float, float] | None:
272
+ """Support of a density such as dbeta(x, 2, 5): [0, 1], not (-10, 10)."""
273
+ if formula.mode != "explicit" or len(formula.variables) != 1:
274
+ return None
275
+ return support_hint(
276
+ getattr(formula, "body", None), formula.variables[0], formula.namespace
277
+ )
278
+
279
+
280
+ def _linspace(lo: float, hi: float, count: int) -> np.ndarray:
281
+ return np.linspace(float(lo), float(hi), int(count))
282
+
283
+
284
+ def _call_formula(formula: Formula, variables: dict[str, np.ndarray]) -> np.ndarray:
285
+ values = evaluate(formula, variables)
286
+ shapes = [np.shape(array) for array in variables.values()]
287
+ target = shapes[0] if shapes else ()
288
+ if values.shape != target:
289
+ try:
290
+ values = np.broadcast_to(values, target).astype(np.float64, copy=True)
291
+ except ValueError as exc:
292
+ raise ExprError(
293
+ "formula result does not match the sampled grid"
294
+ ) from exc
295
+ return np.asarray(values, dtype=np.float64)
296
+
297
+
298
+ def _expand_curve(geom: _Geom, formula: Formula, axes: _Axes, domains: dict) -> _Geom:
299
+ # Sideways ``x = f(y)`` samples the vertical axis. Everything else samples x.
300
+ sample_axis = "y" if axes.computed == "x" else "x"
301
+ (lo, hi), source = _domain_for(geom, sample_axis, domains)
302
+ if source == "default":
303
+ hint = _density_domain(formula)
304
+ if hint is not None:
305
+ (lo, hi), source = hint, "density"
306
+ count = _sample_count(geom, grid=False)
307
+ samples = _linspace(lo, hi, count)
308
+ lo, hi, samples = _narrow_curve(
309
+ formula, axes, samples, lo, hi, source, count
310
+ )
311
+ values = _curve_values(formula, axes, samples)
312
+ if getattr(geom, "n", None) is None:
313
+ samples, values = _refine_curve(formula, axes, samples, values)
314
+ view_axis = "x" if axes.computed == "x" else "y"
315
+ view_lim = _limit_pair(getattr(geom, view_axis + "lim", None), view_axis + "lim")
316
+ view_name = axes.x if view_axis == "x" else axes.y
317
+ kept_s, kept_v, lock, index, note = _clip_series(
318
+ samples, values, view_lim, view_axis, view_name,
319
+ probe=_curve_probe(formula, axes, samples),
320
+ )
321
+ if axes.computed == "x":
322
+ xs, ys = kept_v, kept_s
323
+ else:
324
+ xs, ys = kept_s, kept_v
325
+ frame = pd.DataFrame({"x": np.asarray(xs, dtype=np.float64), "y": np.asarray(ys, dtype=np.float64)})
326
+ groups = _groups_from_index(index)
327
+ if not any(count >= 2 for _start, count in groups):
328
+ raise ExprError("geom_function() needs at least two points on this domain")
329
+ # Sideways ``x = f(y)`` must keep sample order. ``_groups`` tells the
330
+ # encoder not to sort the line and where to break it.
331
+ maker = geom_path if axes.computed == "x" else geom_line
332
+ linewidth = getattr(geom, "linewidth", None)
333
+ out = maker(
334
+ aes(x="x", y="y"),
335
+ linewidth=2.0 if linewidth is None else linewidth,
336
+ color=geom.const_color,
337
+ alpha=geom.alpha,
338
+ )
339
+ out.data_override = frame
340
+ out._groups = groups
341
+ out._replace_mapping = True
342
+ _stamp_formula(out, geom, formula)
343
+ out._axis_labels = {"x": axes.x, "y": axes.y}
344
+ if lock is not None:
345
+ out._axis_lock = {view_axis: lock}
346
+ if note:
347
+ out._notes = [note]
348
+ return out
349
+
350
+
351
+ _REFINE_ROUNDS = 8
352
+ _REFINE_MAX_EXTRA = 4000
353
+
354
+
355
+ def _refine_curve(formula, axes, samples, values):
356
+ """More samples where the curve bends sharply between them.
357
+
358
+ A narrow peak can fall between evenly spaced samples and be drawn short
359
+ (exp(-2000 x^2) topping out at 0.71). Wherever three neighbours bend by
360
+ more than a small share of the curve's height, the two gaps around the
361
+ middle one get a midpoint, and again, up to a few thousand points.
362
+ Poles (non-finite values) are left to the clipping that follows.
363
+ """
364
+ samples = np.asarray(samples, dtype=np.float64)
365
+ values = np.asarray(values, dtype=np.float64)
366
+ finite = np.isfinite(values)
367
+ if finite.sum() < 3:
368
+ return samples, values
369
+ span = float(np.nanmax(values[finite]) - np.nanmin(values[finite]))
370
+ if not np.isfinite(span) or span <= 0:
371
+ return samples, values
372
+ tol = 0.002 * span
373
+ added = 0
374
+ for _ in range(_REFINE_ROUNDS):
375
+ v = values
376
+ left, mid, right = v[:-2], v[1:-1], v[2:]
377
+ bend = np.abs(left - 2.0 * mid + right)
378
+ # Only peaks and valleys: a curve climbing toward a pole bends hard
379
+ # too, and the clipping that follows must see it as it is.
380
+ turning = ((mid >= left) & (mid >= right)) | ((mid <= left) & (mid <= right))
381
+ flagged = np.flatnonzero(np.isfinite(bend) & (bend > tol) & turning) + 1
382
+ if flagged.size == 0:
383
+ break
384
+ gaps = np.unique(np.concatenate([flagged - 1, flagged]))
385
+ gaps = gaps[(gaps >= 0) & (gaps < samples.size - 1)]
386
+ gaps = gaps[np.isfinite(values[gaps]) & np.isfinite(values[gaps + 1])]
387
+ # A big jump across zero is a pole between the samples (tan x).
388
+ jump = (np.sign(values[gaps]) != np.sign(values[gaps + 1])) & (
389
+ np.abs(values[gaps] - values[gaps + 1]) > 0.5 * span
390
+ )
391
+ gaps = gaps[~jump]
392
+ # Gaps already finer than float noise cannot be split usefully.
393
+ width = samples[gaps + 1] - samples[gaps]
394
+ gaps = gaps[width > 1e-12 * max(1.0, float(np.abs(samples).max()))]
395
+ if gaps.size == 0 or added + gaps.size > _REFINE_MAX_EXTRA:
396
+ break
397
+ mids = 0.5 * (samples[gaps] + samples[gaps + 1])
398
+ mid_values = _curve_values(formula, axes, mids)
399
+ samples = np.insert(samples, gaps + 1, mids)
400
+ values = np.insert(values, gaps + 1, mid_values)
401
+ added += gaps.size
402
+ return samples, values
403
+
404
+
405
+ def _curve_values(formula: Formula, axes: _Axes, samples: np.ndarray) -> np.ndarray:
406
+ del axes
407
+ if formula.mode == "callable":
408
+ name = formula.fn_args[0]
409
+ return _call_formula(formula, {name: samples})
410
+ if not formula.variables:
411
+ raw = evaluate(formula, {})
412
+ number = float(np.asarray(raw, dtype=np.float64).reshape(-1)[0])
413
+ return np.full(samples.shape, number, dtype=np.float64)
414
+ return _call_formula(formula, {formula.variables[0]: samples})
415
+
416
+
417
+ def _narrow_curve(
418
+ formula: Formula,
419
+ axes: _Axes,
420
+ samples: np.ndarray,
421
+ lo: float,
422
+ hi: float,
423
+ source: str,
424
+ count: int,
425
+ ) -> tuple[float, float, np.ndarray]:
426
+ """Shrink the default domain when the formula is undefined on most of it."""
427
+ if source != "default":
428
+ return lo, hi, samples
429
+ values = _curve_values(formula, axes, samples)
430
+ finite = np.isfinite(values)
431
+ fraction = float(np.mean(finite)) if finite.size else 0.0
432
+ if fraction >= 0.55 or fraction == 0.0:
433
+ return lo, hi, samples
434
+ good = samples[finite]
435
+ nlo, nhi = float(np.min(good)), float(np.max(good))
436
+ if nhi <= nlo:
437
+ return lo, hi, samples
438
+ narrowed = _linspace(nlo, nhi, count)
439
+ return nlo, nhi, narrowed
440
+
441
+
442
+ def _clip_series(
443
+ samples: np.ndarray,
444
+ values: np.ndarray,
445
+ view_lim: tuple[float, float] | None,
446
+ view_axis: str,
447
+ note_name: str | None = None,
448
+ probe: Callable[..., bool] | None = None,
449
+ ) -> tuple[np.ndarray, np.ndarray, tuple[float, float] | None, np.ndarray, str | None]:
450
+ finite = np.isfinite(samples) & np.isfinite(values)
451
+ if not np.any(finite):
452
+ raise ExprError("geom_function() is undefined everywhere on this domain")
453
+ note: str | None = None
454
+ if view_lim is not None:
455
+ lo, hi = view_lim
456
+ keep = finite & (values >= lo) & (values <= hi)
457
+ lock: tuple[float, float] | None = (lo, hi)
458
+ else:
459
+ lo, hi, blew_up = _robust_window(np.where(finite, values, np.nan), probe)
460
+ # A pole is a thin spike. A piecewise curve (flat, then a parabola)
461
+ # puts a large share of its samples outside that window; keep them.
462
+ outside = finite & ((values < lo) | (values > hi))
463
+ thin = float(np.count_nonzero(outside)) < 0.2 * float(np.count_nonzero(finite))
464
+ if blew_up and thin:
465
+ keep = finite & (values >= lo) & (values <= hi)
466
+ note = _clip_note(note_name or view_axis, lo, hi, param=view_axis)
467
+ else:
468
+ keep = finite
469
+ lock = None
470
+ if not np.any(keep):
471
+ raise ExprError(
472
+ f"geom_function() has no points inside {view_axis}lim="
473
+ f"({view_lim[0]:.6g}, {view_lim[1]:.6g})"
474
+ if view_lim is not None
475
+ else "geom_function() is undefined everywhere on this domain"
476
+ )
477
+ index = np.flatnonzero(keep)
478
+ return samples[index], values[index], lock, index, note
479
+
480
+
481
+ def _groups_from_index(index: np.ndarray) -> list[list[int]]:
482
+ """Break the line wherever clipped samples left a gap."""
483
+ if index.size == 0:
484
+ return []
485
+ groups: list[list[int]] = []
486
+ start = 0
487
+ for position in range(1, int(index.size)):
488
+ if int(index[position]) != int(index[position - 1]) + 1:
489
+ groups.append([start, position - start])
490
+ start = position
491
+ groups.append([start, int(index.size) - start])
492
+ return groups
493
+
494
+
495
+ def _bound(value: float) -> str:
496
+ """About four significant figures, with a Unicode minus."""
497
+ return _number(float(value)).pretty
498
+
499
+
500
+ def _clip_note(name: str, lo: float, hi: float, *, param: str | None = None) -> str:
501
+ """Caption for a pole that was clipped to the bulk of the samples.
502
+
503
+ ``name`` is the variable the reader sees (``t`` on a surface of ``t``).
504
+ ``param`` is the keyword that changes the window (``zlim``).
505
+ """
506
+ flag = param or name
507
+ return (
508
+ f"{name} clipped to [{_bound(lo)}, {_bound(hi)}]; "
509
+ f"pass {flag}lim= to change"
510
+ )
511
+
512
+
513
+ def _stamp_formula(out, geom, formula: Formula) -> None:
514
+ """Legend text, and the symbolic form the tooltip shows above the values."""
515
+ custom = getattr(geom, "label", None)
516
+ if custom:
517
+ plain, segments = split_math(str(custom))
518
+ out._legend_label = plain
519
+ out._legend_math = segments
520
+ out._legend_latex = None
521
+ if formula.mode == "callable":
522
+ if segments and len(segments) == 1 and segments[0].get("latex"):
523
+ out._tip_latex = segments[0]["latex"]
524
+ out._tip_pretty = segments[0]["text"]
525
+ else:
526
+ out._tip_latex = ""
527
+ out._tip_pretty = plain
528
+ else:
529
+ out._tip_latex = formula.caption_latex or formula.latex
530
+ out._tip_pretty = formula.caption_pretty or formula.pretty
531
+ else:
532
+ pretty = formula.pretty or formula.label
533
+ caption = formula.caption_pretty or pretty
534
+ # The title stays the short symbolic formula. The legend keeps the
535
+ # symbols and lists the values, so a Beta density does not expand
536
+ # into ``x^(2 − 1)(1 − x)^(5 − 1)/0.0333333333333``.
537
+ out._title_label = pretty
538
+ out._title_latex = formula.latex or None
539
+ if caption != pretty:
540
+ out._legend_label = caption
541
+ out._legend_latex = formula.caption_latex or formula.latex
542
+ else:
543
+ out._legend_label = formula.legend_pretty or pretty
544
+ out._legend_latex = formula.legend_latex or formula.latex or None
545
+ out._legend_math = None
546
+ out._tip_latex = formula.caption_latex or formula.latex or ""
547
+ out._tip_pretty = caption
548
+ out._is_formula = True
549
+ out._formula_primary = True
550
+
551
+
552
+ def _robust_window(
553
+ values: np.ndarray, probe: Callable[..., bool] | None = None
554
+ ) -> tuple[float, float, bool]:
555
+ """Return ``(lo, hi, blew_up)`` around the bulk of ``values``.
556
+
557
+ Only a side that actually leaves the bulk is pulled in. A density that
558
+ stays non-negative keeps its own minimum instead of a negative fence.
559
+
560
+ ``probe(index, value, sign, centre)`` confirms that the extreme at
561
+ ``values[index]`` really grows without bound (see ``_keeps_growing``).
562
+ Without it, anything past the fence counts as a blow-up.
563
+ """
564
+ values = np.asarray(values, dtype=np.float64)
565
+ where = np.flatnonzero(np.isfinite(values))
566
+ finite = values[where]
567
+ med = float(np.median(finite))
568
+ mad = float(np.median(np.abs(finite - med)))
569
+ scale = max(mad * 1.4826, 1e-9)
570
+ fence_lo = med - 8.0 * scale
571
+ fence_hi = med + 8.0 * scale
572
+ full_lo = float(np.min(finite))
573
+ full_hi = float(np.max(finite))
574
+ blew_lo = full_lo < fence_lo - 1e-8
575
+ blew_hi = full_hi > fence_hi + 1e-8
576
+ # A steep but finite curve (Beta(5, 1) = 5x^4 near x = 1) also leaves
577
+ # the fence. Keep its true extreme unless zooming in shows a pole.
578
+ if probe is not None:
579
+ if blew_hi:
580
+ index = int(where[int(np.argmax(finite))])
581
+ blew_hi = probe(index, full_hi, 1.0, med)
582
+ if blew_lo:
583
+ index = int(where[int(np.argmin(finite))])
584
+ blew_lo = probe(index, full_lo, -1.0, med)
585
+ lo = fence_lo if blew_lo else full_lo
586
+ hi = fence_hi if blew_hi else full_hi
587
+ if hi <= lo:
588
+ hi = lo + 1.0
589
+ return lo, hi, blew_lo or blew_hi
590
+
591
+
592
+ _ZOOM_POINTS = 41
593
+ _ZOOM_LEVELS = 2
594
+ _ZOOM_GROWTH = 1.5
595
+ _ZOOM_HUGE = 1e6
596
+
597
+
598
+ def _keeps_growing(
599
+ evaluate: Callable[..., np.ndarray],
600
+ box: list[tuple[float, float]],
601
+ start: float,
602
+ sign: float,
603
+ centre: float,
604
+ ) -> bool:
605
+ """True when the extreme inside ``box`` grows on every zoom.
606
+
607
+ ``box`` holds one ``(lo, hi)`` interval per axis, the neighbours of the
608
+ extreme sample. ``evaluate(*axes)`` returns an array whose axis ``k``
609
+ follows ``box[k]``. A pole (``1/x``, ``tan x``) moves further from the
610
+ bulk each time the samples close in on it. A finite maximum, at an edge
611
+ or in a narrow peak between samples, settles after the first zoom.
612
+ """
613
+ distance = sign * (start - centre)
614
+ if distance <= 0.0:
615
+ return True
616
+ first = distance
617
+ for _level in range(_ZOOM_LEVELS):
618
+ axes = [np.linspace(lo, hi, _ZOOM_POINTS) for lo, hi in box]
619
+ try:
620
+ with np.errstate(all="ignore"):
621
+ toward = sign * np.asarray(evaluate(*axes), dtype=np.float64)
622
+ except Exception:
623
+ return True # cannot tell; keep the old behaviour
624
+ if np.any(np.isposinf(toward)):
625
+ return True # landed on the pole itself
626
+ ok = np.isfinite(toward)
627
+ if not np.any(ok):
628
+ return True
629
+ flat = np.where(ok, toward, -np.inf)
630
+ best = int(np.argmax(flat))
631
+ reached = float(flat.reshape(-1)[best]) - sign * centre
632
+ ratio = reached / distance
633
+ if ratio < 0.5:
634
+ # The sampled extreme vanished when resampled: a rounding spike
635
+ # beside a singularity (x*y/(x^2 - y^2) on the diagonal).
636
+ return True
637
+ if reached > first * _ZOOM_HUGE:
638
+ # Float precision stops the next zoom from closing in further.
639
+ return True
640
+ position = np.unravel_index(best, flat.shape)
641
+ if ratio < _ZOOM_GROWTH:
642
+ # Settled. A true maximum is continuous: the zoom points right
643
+ # beside it are nearly as high. An isolated point is rounding
644
+ # noise at a singularity, which should still be clipped.
645
+ near = tuple(
646
+ slice(max(int(i) - 1, 0), int(i) + 2) for i in position
647
+ )
648
+ around = flat[near].copy()
649
+ around[tuple(int(i) - s.start for i, s in zip(position, near))] = -np.inf
650
+ beside = float(np.max(around)) - sign * centre
651
+ return bool(beside < 0.5 * reached)
652
+ distance = reached
653
+ box = [
654
+ (
655
+ float(axis[max(int(i) - 1, 0)]),
656
+ float(axis[min(int(i) + 1, _ZOOM_POINTS - 1)]),
657
+ )
658
+ for axis, i in zip(axes, position)
659
+ ]
660
+ return True
661
+
662
+
663
+ def _neighbours(axis: np.ndarray, i: int) -> tuple[float, float]:
664
+ return float(axis[max(i - 1, 0)]), float(axis[min(i + 1, axis.size - 1)])
665
+
666
+
667
+ def _curve_probe(formula: Formula, axes: _Axes, samples: np.ndarray):
668
+ def evaluate(xs: np.ndarray) -> np.ndarray:
669
+ return _curve_values(formula, axes, xs)
670
+
671
+ def probe(index: int, value: float, sign: float, centre: float) -> bool:
672
+ box = [_neighbours(samples, index)]
673
+ return _keeps_growing(evaluate, box, value, sign, centre)
674
+
675
+ return probe
676
+
677
+
678
+ def _surface_probe(formula: Formula, axes: _Axes, xs: np.ndarray, ys: np.ndarray):
679
+ """Probe for a flat index into a ``(len(ys), len(xs))`` grid."""
680
+
681
+ def evaluate(x_axis: np.ndarray, y_axis: np.ndarray) -> np.ndarray:
682
+ return _surface_values(formula, axes, x_axis, y_axis).T
683
+
684
+ def probe(index: int, value: float, sign: float, centre: float) -> bool:
685
+ row, col = divmod(index, xs.size)
686
+ box = [_neighbours(xs, col), _neighbours(ys, row)]
687
+ return _keeps_growing(evaluate, box, value, sign, centre)
688
+
689
+ return probe
690
+
691
+
692
+ def _expand_surface(
693
+ geom: _Geom, formula: Formula, axes: _Axes, domains: dict
694
+ ) -> _Geom:
695
+ count = _sample_count(geom, grid=True)
696
+ (xlo, xhi), x_source = _domain_for(geom, "x", domains)
697
+ (ylo, yhi), y_source = _domain_for(geom, "y", domains)
698
+ xs = _linspace(xlo, xhi, count)
699
+ ys = _linspace(ylo, yhi, count)
700
+ zz = _surface_values(formula, axes, xs, ys)
701
+ xs, ys, zz = _narrow_surface(
702
+ formula,
703
+ axes,
704
+ xs,
705
+ ys,
706
+ zz,
707
+ x_source == "default" and _limit_pair(geom.xlim, "xlim") is None,
708
+ y_source == "default" and _limit_pair(geom.ylim, "ylim") is None,
709
+ count,
710
+ )
711
+ zz, lock, note = _clip_grid(
712
+ zz, _limit_pair(geom.zlim, "zlim"), axes.z or "z",
713
+ probe=_surface_probe(formula, axes, xs, ys),
714
+ )
715
+ xx, yy = np.meshgrid(xs, ys)
716
+ frame = pd.DataFrame(
717
+ {
718
+ "x": xx.ravel(),
719
+ "y": yy.ravel(),
720
+ "z": zz.ravel(),
721
+ }
722
+ )
723
+ # No colour of your own: colour by height, so the shape reads in print.
724
+ by_height = geom.const_color is None
725
+ if by_height:
726
+ frame["height"] = frame["z"]
727
+ vertices, indices, nx, ny = regular_grid_mesh(
728
+ frame, "x", "y", "z", ccol="height" if by_height else None
729
+ )
730
+ colour_col = "colour" if by_height and "colour" in vertices.columns else None
731
+ out = _Geom(
732
+ aes(x="x", y="y", z="z", colour=colour_col),
733
+ color=geom.const_color,
734
+ alpha=geom.alpha if geom.alpha is not None else 0.95,
735
+ )
736
+ out.kind = "surface"
737
+ if colour_col:
738
+ out._default_ramp = "viridis"
739
+ out._colour_title = axes.z or "z"
740
+ out.data_override = vertices
741
+ out.const_color = geom.const_color
742
+ out.alpha = geom.alpha if geom.alpha is not None else 0.95
743
+ out.wireframe = bool(getattr(geom, "wireframe", False))
744
+ out._indices = indices
745
+ out._nx = nx
746
+ out._ny = ny
747
+ out._replace_mapping = True
748
+ _stamp_formula(out, geom, formula)
749
+ out._axis_labels = {"x": axes.x, "y": axes.y, "z": axes.z or "z"}
750
+ out._function_surface = True
751
+ if lock is not None:
752
+ out._axis_lock = {"z": lock}
753
+ if note:
754
+ out._notes = [note]
755
+ return out
756
+
757
+
758
+ def _surface_values(
759
+ formula: Formula, axes: _Axes, xs: np.ndarray, ys: np.ndarray
760
+ ) -> np.ndarray:
761
+ xx, yy = np.meshgrid(xs, ys)
762
+ if formula.mode == "callable":
763
+ return _call_formula(
764
+ formula,
765
+ {formula.fn_args[0]: xx, formula.fn_args[1]: yy},
766
+ )
767
+ return _call_formula(formula, {axes.x: xx, axes.y: yy})
768
+
769
+
770
+ def _narrow_surface(
771
+ formula: Formula,
772
+ axes: _Axes,
773
+ xs: np.ndarray,
774
+ ys: np.ndarray,
775
+ zz: np.ndarray,
776
+ narrow_x: bool,
777
+ narrow_y: bool,
778
+ count: int,
779
+ ) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
780
+ finite = np.isfinite(zz)
781
+ fraction = float(np.mean(finite)) if finite.size else 0.0
782
+ if fraction >= 0.55 or fraction == 0.0 or not (narrow_x or narrow_y):
783
+ return xs, ys, zz
784
+ rows = np.any(finite, axis=1)
785
+ cols = np.any(finite, axis=0)
786
+ if narrow_y and np.any(rows):
787
+ ys = _linspace(float(ys[rows][0]), float(ys[rows][-1]), count)
788
+ if narrow_x and np.any(cols):
789
+ xs = _linspace(float(xs[cols][0]), float(xs[cols][-1]), count)
790
+ return xs, ys, _surface_values(formula, axes, xs, ys)
791
+
792
+
793
+ def _clip_grid(
794
+ zz: np.ndarray,
795
+ zlim: tuple[float, float] | None,
796
+ note_name: str = "z",
797
+ probe: Callable[..., bool] | None = None,
798
+ ) -> tuple[np.ndarray, tuple[float, float] | None, str | None]:
799
+ finite = zz[np.isfinite(zz)]
800
+ if finite.size == 0:
801
+ raise ExprError("geom_function() is undefined everywhere on this domain")
802
+ if zlim is not None:
803
+ lo, hi = zlim
804
+ clipped = np.clip(np.where(np.isfinite(zz), zz, lo), lo, hi)
805
+ return clipped, (lo, hi), None
806
+ lo, hi, blew_up = _robust_window(zz.ravel(), probe)
807
+ if not blew_up:
808
+ filled = np.where(np.isfinite(zz), zz, float(np.median(finite)))
809
+ return filled, None, None
810
+ filled = np.where(np.isfinite(zz), zz, lo)
811
+ filled = np.clip(filled, lo, hi)
812
+ return filled, (lo, hi), _clip_note(note_name, lo, hi, param="z")
813
+
814
+
815
+ def _expand_implicit(
816
+ geom: _Geom, formula: Formula, axes: _Axes, domains: dict
817
+ ) -> _Geom:
818
+ count = _sample_count(geom, grid=True)
819
+ (xlo, xhi), _x_source = _domain_for(geom, "x", domains)
820
+ (ylo, yhi), _y_source = _domain_for(geom, "y", domains)
821
+ xs = _linspace(xlo, xhi, count)
822
+ ys = _linspace(ylo, yhi, count)
823
+ xx, yy = np.meshgrid(xs, ys)
824
+ field = _call_formula(formula, {axes.x: xx, axes.y: yy})
825
+
826
+ def sample(xx_fine: np.ndarray, yy_fine: np.ndarray) -> np.ndarray:
827
+ return _call_formula(formula, {axes.x: xx_fine, axes.y: yy_fine})
828
+
829
+ # The base grid stays at ``n`` (default 80). Only cells the contour
830
+ # crosses are subdivided, so a small loop stays smooth beside a long
831
+ # curve. A full-window refit cannot separate those two pieces. If the
832
+ # finer grid would be too large, keep the coarse contour.
833
+ polylines = _refine_active_cells(xs, ys, field, 0.0, sample)
834
+ if polylines is None:
835
+ polylines = _contour_lines(xs, ys, field, 0.0)
836
+ if not polylines:
837
+ raise ExprError(
838
+ "geom_function() found no curve where the equation is zero "
839
+ "on this domain. Try a wider xlim= and ylim="
840
+ )
841
+ rows_x: list[float] = []
842
+ rows_y: list[float] = []
843
+ groups: list[list[int]] = []
844
+ for poly in polylines:
845
+ if len(poly) < 2:
846
+ continue
847
+ start = len(rows_x)
848
+ for x_val, y_val in poly:
849
+ rows_x.append(x_val)
850
+ rows_y.append(y_val)
851
+ groups.append([start, len(rows_x) - start])
852
+ if not groups:
853
+ raise ExprError(
854
+ "geom_function() found no curve where the equation is zero "
855
+ "on this domain. Try a wider xlim= and ylim="
856
+ )
857
+ frame = pd.DataFrame({"x": rows_x, "y": rows_y})
858
+ linewidth = getattr(geom, "linewidth", None)
859
+ out = geom_path(
860
+ aes(x="x", y="y"),
861
+ linewidth=2.0 if linewidth is None else linewidth,
862
+ color=geom.const_color,
863
+ alpha=geom.alpha,
864
+ )
865
+ out.data_override = frame
866
+ out.sort_x = False
867
+ out._groups = groups
868
+ out._replace_mapping = True
869
+ out._implicit = True
870
+ _stamp_formula(out, geom, formula)
871
+ out._axis_labels = {"x": axes.x, "y": axes.y}
872
+ return out
873
+
874
+
875
+ _MISSING = object()
876
+
877
+
878
+ @contextmanager
879
+ def _bound_params(formula: Formula, params: dict[str, float]):
880
+ """Inject sweep values for one frame, then restore the namespace."""
881
+ saved: list[tuple[str, Any]] = []
882
+ try:
883
+ for name, value in params.items():
884
+ saved.append((name, formula.namespace.get(name, _MISSING)))
885
+ formula.namespace[name] = float(value)
886
+ yield
887
+ finally:
888
+ for name, old in saved:
889
+ if old is _MISSING:
890
+ formula.namespace.pop(name, None)
891
+ else:
892
+ formula.namespace[name] = old
893
+
894
+
895
+ def _sweep(transition: Any) -> list[dict[str, float]]:
896
+ """One shared step for every parameter, from lo to hi across ``frames``."""
897
+ count = int(transition.frames)
898
+ weights = np.linspace(0.0, 1.0, count)
899
+ steps: list[dict[str, float]] = []
900
+ for weight in weights:
901
+ t = float(weight)
902
+ steps.append(
903
+ {
904
+ name: float(lo + t * (hi - lo))
905
+ for name, (lo, hi) in transition.ranges.items()
906
+ }
907
+ )
908
+ return steps
909
+
910
+
911
+ # nSamples * nFrames. A 501-point curve at 25x25 fits; a default surface
912
+ # (80x80) times two 25-step sliders does not. Raise instead of thinning.
913
+ _SLIDER_CELL_CAP = 500_000
914
+
915
+
916
+ def _slider_frame_count(source: Any) -> int:
917
+ if getattr(source, "kind", None) != "slider":
918
+ return 0
919
+ n = 1
920
+ for _name in source.ranges:
921
+ n *= int(source.steps)
922
+ return n
923
+
924
+
925
+ def _guard_slider(source: Any, n_samples: int, what: str) -> None:
926
+ n_frames = _slider_frame_count(source)
927
+ if not n_frames:
928
+ return
929
+ cells = int(n_samples) * n_frames
930
+ if cells <= _SLIDER_CELL_CAP:
931
+ return
932
+ raise ValueError(
933
+ f"slider() would sample {n_frames} frames of {n_samples} {what} "
934
+ f"({cells} values). Pass a smaller steps= or n=."
935
+ )
936
+
937
+
938
+ def _parameter_steps(source: Any) -> list[dict[str, float]]:
939
+ """Frame parameters. A transition locksteps; a slider is the full grid.
940
+
941
+ Grid order is C-order, last keyword fastest, matching
942
+ ``np.meshgrid(..., indexing='ij')`` then ravel.
943
+ """
944
+ if getattr(source, "kind", None) != "slider":
945
+ return _sweep(source)
946
+ names = list(source.ranges)
947
+ axes = [
948
+ np.linspace(float(lo), float(hi), int(source.steps))
949
+ for lo, hi in source.ranges.values()
950
+ ]
951
+ grids = np.meshgrid(*axes, indexing="ij")
952
+ flat = [np.asarray(g, dtype=np.float64).ravel() for g in grids]
953
+ n = int(flat[0].size)
954
+ return [
955
+ {name: float(flat[k][i]) for k, name in enumerate(names)}
956
+ for i in range(n)
957
+ ]
958
+
959
+
960
+ def _static_col(source: Any) -> int:
961
+ """Column shown with JavaScript off.
962
+
963
+ Sliders open with every thumb at the low end. A transition keeps the
964
+ last frame, which is the frame a paused chart already showed.
965
+ """
966
+ if getattr(source, "kind", None) == "slider":
967
+ return 0
968
+ return -1
969
+
970
+
971
+ def _shared_frame_window(
972
+ mat: np.ndarray,
973
+ probe: Callable[..., bool] | None = None,
974
+ ) -> tuple[float, float, bool]:
975
+ """One clip window from every frame, without letting quiet frames shrink it.
976
+
977
+ A frame that stays inside its own robust window contributes its true
978
+ min and max. A frame with a pole contributes only that robust window.
979
+ Quiet frames (``a = 0``) would otherwise pull a pooled median toward
980
+ zero and clip a wave that is perfectly finite on its own.
981
+ """
982
+ healthy_lo = float("inf")
983
+ healthy_hi = -float("inf")
984
+ robust_lo = float("inf")
985
+ robust_hi = -float("inf")
986
+ any_blow = False
987
+ any_healthy = False
988
+ for col in range(mat.shape[1]):
989
+ column = mat[:, col]
990
+ values = column[np.isfinite(column)]
991
+ if values.size == 0:
992
+ continue
993
+ frame_probe = None if probe is None else partial(probe, col)
994
+ lo, hi, blew_up = _robust_window(column, frame_probe)
995
+ if blew_up:
996
+ any_blow = True
997
+ robust_lo = min(robust_lo, lo)
998
+ robust_hi = max(robust_hi, hi)
999
+ else:
1000
+ any_healthy = True
1001
+ healthy_lo = min(healthy_lo, float(np.min(values)))
1002
+ healthy_hi = max(healthy_hi, float(np.max(values)))
1003
+ if not any_blow:
1004
+ return 0.0, 0.0, False
1005
+ if any_healthy:
1006
+ lo = min(robust_lo, healthy_lo)
1007
+ hi = max(robust_hi, healthy_hi)
1008
+ else:
1009
+ lo, hi = robust_lo, robust_hi
1010
+ if hi <= lo:
1011
+ hi = lo + 1.0
1012
+ return lo, hi, True
1013
+
1014
+
1015
+ def _clip_matrix(
1016
+ mat: np.ndarray,
1017
+ view_lim: tuple[float, float] | None,
1018
+ view_axis: str,
1019
+ note_name: str | None,
1020
+ probe: Callable[..., bool] | None = None,
1021
+ ) -> tuple[np.ndarray, tuple[float, float] | None, str | None]:
1022
+ """One window for every frame. Vertices stay; poles are clipped, not dropped."""
1023
+ finite = mat[np.isfinite(mat)]
1024
+ if finite.size == 0:
1025
+ raise ExprError("geom_function() is undefined everywhere on this domain")
1026
+ if view_lim is not None:
1027
+ lo, hi = view_lim
1028
+ filled = np.where(np.isfinite(mat), mat, lo)
1029
+ return np.clip(filled, lo, hi), (lo, hi), None
1030
+ lo, hi, blew_up = _shared_frame_window(mat, probe)
1031
+ if blew_up:
1032
+ filled = np.where(np.isfinite(mat), mat, lo)
1033
+ return (
1034
+ np.clip(filled, lo, hi),
1035
+ (lo, hi),
1036
+ _clip_note(note_name or view_axis, lo, hi, param=view_axis),
1037
+ )
1038
+ fill = float(np.median(finite))
1039
+ return np.where(np.isfinite(mat), mat, fill), None, None
1040
+
1041
+
1042
+ def _curve_matrix(
1043
+ formula: Formula, samples: np.ndarray, steps: list[dict[str, float]]
1044
+ ) -> np.ndarray:
1045
+ columns = []
1046
+ for params in steps:
1047
+ with _bound_params(formula, params):
1048
+ columns.append(_curve_values(formula, None, samples))
1049
+ return np.column_stack(columns)
1050
+
1051
+
1052
+ def _expand_animated(geom, formula, axes, domains, transition) -> _Geom:
1053
+ if axes.kind == "surface":
1054
+ return _expand_surface_anim(geom, formula, axes, domains, transition)
1055
+ if axes.kind == "implicit":
1056
+ return _expand_implicit_anim(geom, formula, axes, domains, transition)
1057
+ return _expand_curve_anim(geom, formula, axes, domains, transition)
1058
+
1059
+
1060
+ def _expand_curve_anim(geom, formula, axes, domains, transition) -> _Geom:
1061
+ # Same samples on every frame, so the line can tween vertex for vertex.
1062
+ sample_axis = "y" if axes.computed == "x" else "x"
1063
+ (lo, hi), source = _domain_for(geom, sample_axis, domains)
1064
+ count = _sample_count(geom, grid=False)
1065
+ _guard_slider(transition, count, "curve samples")
1066
+ steps = _parameter_steps(transition)
1067
+ if source == "default":
1068
+ # dbeta(x, a, b) under a slider: every frame's support, once.
1069
+ spans = []
1070
+ for params in steps:
1071
+ with _bound_params(formula, params):
1072
+ spans.append(_density_domain(formula))
1073
+ if spans and all(span is not None for span in spans):
1074
+ lo = min(span[0] for span in spans)
1075
+ hi = max(span[1] for span in spans)
1076
+ source = "density"
1077
+ samples = _linspace(lo, hi, count)
1078
+ mat = _curve_matrix(formula, samples, steps)
1079
+ if source == "default":
1080
+ finite = np.isfinite(mat)
1081
+ fraction = float(np.mean(finite)) if finite.size else 0.0
1082
+ if 0.0 < fraction < 0.55:
1083
+ good = np.any(finite, axis=1)
1084
+ if np.any(good):
1085
+ nlo = float(np.min(samples[good]))
1086
+ nhi = float(np.max(samples[good]))
1087
+ if nhi > nlo:
1088
+ samples = _linspace(nlo, nhi, count)
1089
+ mat = _curve_matrix(formula, samples, steps)
1090
+ view_axis = "x" if axes.computed == "x" else "y"
1091
+ view_lim = _limit_pair(getattr(geom, view_axis + "lim", None), view_axis + "lim")
1092
+ view_name = axes.x if view_axis == "x" else axes.y
1093
+ def frame_probe(col: int, index: int, value: float, sign: float, centre: float) -> bool:
1094
+ with _bound_params(formula, steps[col]):
1095
+ return _curve_probe(formula, axes, samples)(index, value, sign, centre)
1096
+
1097
+ mat, lock, note = _clip_matrix(mat, view_lim, view_axis, view_name, frame_probe)
1098
+ n_frames = mat.shape[1]
1099
+ repeated = np.repeat(samples[:, None], n_frames, axis=1)
1100
+ if axes.computed == "x":
1101
+ x_mat, y_mat = mat, repeated
1102
+ else:
1103
+ x_mat, y_mat = repeated, mat
1104
+ shown = _static_col(transition)
1105
+ frame = pd.DataFrame(
1106
+ {
1107
+ "x": np.asarray(x_mat[:, shown], dtype=np.float64),
1108
+ "y": np.asarray(y_mat[:, shown], dtype=np.float64),
1109
+ }
1110
+ )
1111
+ maker = geom_path if axes.computed == "x" else geom_line
1112
+ linewidth = getattr(geom, "linewidth", None)
1113
+ out = maker(
1114
+ aes(x="x", y="y"),
1115
+ linewidth=2.0 if linewidth is None else linewidth,
1116
+ color=geom.const_color,
1117
+ alpha=geom.alpha,
1118
+ )
1119
+ out.data_override = frame
1120
+ out._groups = [[0, int(samples.size)]]
1121
+ out._replace_mapping = True
1122
+ _stamp_formula(out, geom, formula)
1123
+ out._axis_labels = {"x": axes.x, "y": axes.y}
1124
+ if lock is not None:
1125
+ out._axis_lock = {view_axis: lock}
1126
+ if note:
1127
+ out._notes = [note]
1128
+ anim = {
1129
+ "mode": "tween",
1130
+ "channels": {"x": x_mat, "y": y_mat},
1131
+ }
1132
+ if shown == 0:
1133
+ anim["static_col"] = 0
1134
+ out._anim = anim
1135
+ return out
1136
+
1137
+
1138
+ def _expand_surface_anim(geom, formula, axes, domains, transition) -> _Geom:
1139
+ count = _sample_count(geom, grid=True)
1140
+ _guard_slider(transition, count * count, "surface vertices")
1141
+ (xlo, xhi), x_source = _domain_for(geom, "x", domains)
1142
+ (ylo, yhi), y_source = _domain_for(geom, "y", domains)
1143
+ xs = _linspace(xlo, xhi, count)
1144
+ ys = _linspace(ylo, yhi, count)
1145
+ steps = _parameter_steps(transition)
1146
+
1147
+ def grids(x_axis: np.ndarray, y_axis: np.ndarray) -> np.ndarray:
1148
+ layers = []
1149
+ for params in steps:
1150
+ with _bound_params(formula, params):
1151
+ layers.append(_surface_values(formula, axes, x_axis, y_axis))
1152
+ return np.stack(layers, axis=-1)
1153
+
1154
+ stack = grids(xs, ys)
1155
+ finite = np.isfinite(stack)
1156
+ fraction = float(np.mean(finite)) if finite.size else 0.0
1157
+ narrow_x = x_source == "default"
1158
+ narrow_y = y_source == "default"
1159
+ if 0.0 < fraction < 0.55 and (narrow_x or narrow_y):
1160
+ any_finite = np.any(finite, axis=-1)
1161
+ rows = np.any(any_finite, axis=1)
1162
+ cols = np.any(any_finite, axis=0)
1163
+ if narrow_y and np.any(rows):
1164
+ ys = _linspace(float(ys[rows][0]), float(ys[rows][-1]), count)
1165
+ if narrow_x and np.any(cols):
1166
+ xs = _linspace(float(xs[cols][0]), float(xs[cols][-1]), count)
1167
+ stack = grids(xs, ys)
1168
+ zlim = _limit_pair(geom.zlim, "zlim")
1169
+ flat = stack.reshape(-1, stack.shape[-1])
1170
+ def frame_probe(col: int, index: int, value: float, sign: float, centre: float) -> bool:
1171
+ with _bound_params(formula, steps[col]):
1172
+ return _surface_probe(formula, axes, xs, ys)(index, value, sign, centre)
1173
+
1174
+ z_mat, lock, note = _clip_matrix(flat, zlim, "z", axes.z or "z", frame_probe)
1175
+ xx, yy = np.meshgrid(xs, ys)
1176
+ n_frames = z_mat.shape[1]
1177
+ shown = _static_col(transition)
1178
+ last = pd.DataFrame(
1179
+ {
1180
+ "x": xx.ravel(),
1181
+ "y": yy.ravel(),
1182
+ "z": np.asarray(z_mat[:, shown], dtype=np.float64),
1183
+ }
1184
+ )
1185
+ vertices, indices, nx, ny = regular_grid_mesh(last, "x", "y", "z")
1186
+ vx = vertices["x"].to_numpy(dtype=np.float64)
1187
+ vy = vertices["y"].to_numpy(dtype=np.float64)
1188
+ if np.allclose(vx, xx.ravel()) and np.allclose(vy, yy.ravel()):
1189
+ src = np.arange(xx.size)
1190
+ else:
1191
+ ix = np.clip(np.searchsorted(xs, vx), 0, len(xs) - 1)
1192
+ iy = np.clip(np.searchsorted(ys, vy), 0, len(ys) - 1)
1193
+ src = iy * len(xs) + ix
1194
+ x_mat = np.repeat(xx.ravel()[:, None], n_frames, axis=1)[src]
1195
+ y_mat = np.repeat(yy.ravel()[:, None], n_frames, axis=1)[src]
1196
+ z_ordered = z_mat[src]
1197
+ out = _Geom(
1198
+ aes(x="x", y="y", z="z"),
1199
+ color=geom.const_color,
1200
+ alpha=geom.alpha if geom.alpha is not None else 0.95,
1201
+ )
1202
+ out.kind = "surface"
1203
+ out.data_override = vertices
1204
+ out.const_color = geom.const_color
1205
+ out.alpha = geom.alpha if geom.alpha is not None else 0.95
1206
+ out.wireframe = bool(getattr(geom, "wireframe", False))
1207
+ out._indices = indices
1208
+ out._nx = nx
1209
+ out._ny = ny
1210
+ out._replace_mapping = True
1211
+ _stamp_formula(out, geom, formula)
1212
+ out._axis_labels = {"x": axes.x, "y": axes.y, "z": axes.z or "z"}
1213
+ out._function_surface = True
1214
+ if lock is not None:
1215
+ out._axis_lock = {"z": lock}
1216
+ if note:
1217
+ out._notes = [note]
1218
+ anim = {
1219
+ "mode": "tween",
1220
+ "channels": {"x": x_mat, "y": y_mat, "z": z_ordered},
1221
+ }
1222
+ if shown == 0:
1223
+ anim["static_col"] = 0
1224
+ out._anim = anim
1225
+ return out
1226
+
1227
+
1228
+ def _rows_from_polylines(polylines) -> tuple[np.ndarray, np.ndarray, list]:
1229
+ rows_x: list[float] = []
1230
+ rows_y: list[float] = []
1231
+ groups: list[list[int]] = []
1232
+ for poly in polylines or []:
1233
+ if len(poly) < 2:
1234
+ continue
1235
+ start = len(rows_x)
1236
+ for x_val, y_val in poly:
1237
+ rows_x.append(float(x_val))
1238
+ rows_y.append(float(y_val))
1239
+ groups.append([start, len(rows_x) - start])
1240
+ return (
1241
+ np.asarray(rows_x, dtype=np.float64),
1242
+ np.asarray(rows_y, dtype=np.float64),
1243
+ groups,
1244
+ )
1245
+
1246
+
1247
+ def _expand_implicit_anim(geom, formula, axes, domains, transition) -> _Geom:
1248
+ # Vertex counts change with the parameter, so frames are shown as-is.
1249
+ count = _sample_count(geom, grid=True)
1250
+ _guard_slider(transition, count * count, "grid samples")
1251
+ (xlo, xhi), _x_source = _domain_for(geom, "x", domains)
1252
+ (ylo, yhi), _y_source = _domain_for(geom, "y", domains)
1253
+ xs = _linspace(xlo, xhi, count)
1254
+ ys = _linspace(ylo, yhi, count)
1255
+ steps = _parameter_steps(transition)
1256
+ frames: list[dict] = []
1257
+ static = None
1258
+ static_i = 0
1259
+ # A slider opens at the low end, so the fallback curve is the first
1260
+ # non-empty frame. A transition keeps the last one.
1261
+ keep_first = getattr(transition, "kind", None) == "slider"
1262
+ for index, params in enumerate(steps):
1263
+
1264
+ def sample(xx_fine, yy_fine, params=params):
1265
+ with _bound_params(formula, params):
1266
+ return _call_formula(formula, {axes.x: xx_fine, axes.y: yy_fine})
1267
+
1268
+ with _bound_params(formula, params):
1269
+ xx, yy = np.meshgrid(xs, ys)
1270
+ field = _call_formula(formula, {axes.x: xx, axes.y: yy})
1271
+ polylines = _refine_active_cells(xs, ys, field, 0.0, sample)
1272
+ if polylines is None:
1273
+ polylines = _contour_lines(xs, ys, field, 0.0)
1274
+ rows_x, rows_y, groups = _rows_from_polylines(polylines)
1275
+ frame = {"x": rows_x, "y": rows_y, "groups": groups}
1276
+ frames.append(frame)
1277
+ if rows_x.size >= 2 and (static is None or not keep_first):
1278
+ static = frame
1279
+ static_i = index
1280
+ if static is None:
1281
+ raise ExprError(
1282
+ "geom_function() found no curve where the equation is zero "
1283
+ "on this domain. Try a wider xlim= and ylim="
1284
+ )
1285
+ drawn = pd.DataFrame({"x": static["x"], "y": static["y"]})
1286
+ linewidth = getattr(geom, "linewidth", None)
1287
+ out = geom_path(
1288
+ aes(x="x", y="y"),
1289
+ linewidth=2.0 if linewidth is None else linewidth,
1290
+ color=geom.const_color,
1291
+ alpha=geom.alpha,
1292
+ )
1293
+ out.data_override = drawn
1294
+ out.sort_x = False
1295
+ out._groups = static["groups"]
1296
+ out._replace_mapping = True
1297
+ out._implicit = True
1298
+ _stamp_formula(out, geom, formula)
1299
+ out._axis_labels = {"x": axes.x, "y": axes.y}
1300
+ out._anim = {"mode": "step", "frames": frames, "static": int(static_i)}
1301
+ return out