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/special.py ADDED
@@ -0,0 +1,407 @@
1
+ """Special functions and probability densities a formula may call.
2
+
3
+ NumPy has no gamma or erf, so these wrap the ``math`` module element-wise.
4
+ Names follow R (``dnorm``, ``pnorm``, ``qnorm``), which ggplot users know::
5
+
6
+ geom_function("y = dbeta(x, 2, 5)") + area(0.2, 0.5)
7
+ geom_function("y = dnorm(x, mu, sigma)", mu=100, sigma=15)
8
+
9
+ A density is 0 outside its support, as in R. ``support()`` gives the range
10
+ worth drawing, so ``dbeta`` opens on [0, 1] instead of the default (-10, 10).
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import ast
16
+ import math
17
+ from typing import Any, Callable
18
+
19
+ import numpy as np
20
+
21
+
22
+ def _elementwise(fn: Callable[[float], float]) -> Callable[..., np.ndarray]:
23
+ def safe(*args: float) -> float:
24
+ try:
25
+ return float(fn(*args))
26
+ except (ValueError, OverflowError, ZeroDivisionError):
27
+ return math.nan
28
+
29
+ vec = np.vectorize(safe, otypes=[np.float64])
30
+
31
+ def apply(*args: Any) -> Any:
32
+ out = vec(*args)
33
+ return out if np.ndim(out) else float(out)
34
+
35
+ apply.__name__ = getattr(fn, "__name__", "special")
36
+ return apply
37
+
38
+
39
+ def _gamma(x: float) -> float:
40
+ if x <= 0 and float(x).is_integer():
41
+ return math.inf # poles at 0, -1, -2, ...
42
+ return math.gamma(x)
43
+
44
+
45
+ gamma = _elementwise(_gamma)
46
+ lgamma = _elementwise(math.lgamma)
47
+ erf = _elementwise(math.erf)
48
+ erfc = _elementwise(math.erfc)
49
+
50
+
51
+ def _lbeta(a, b):
52
+ return lgamma(a) + lgamma(b) - lgamma(np.add(a, b))
53
+
54
+
55
+ def beta(a, b):
56
+ """Euler's Beta function B(a, b) for a, b > 0."""
57
+ return np.exp(_lbeta(a, b))
58
+
59
+
60
+ def _f(x) -> np.ndarray:
61
+ return np.asarray(x, dtype=np.float64)
62
+
63
+
64
+ def dnorm(x, mean=0.0, sd=1.0):
65
+ x, mean, sd = _f(x), _f(mean), _f(sd)
66
+ with np.errstate(all="ignore"):
67
+ z = (x - mean) / sd
68
+ return np.where(sd > 0, np.exp(-0.5 * z * z) / (sd * math.sqrt(2.0 * math.pi)), np.nan)
69
+
70
+
71
+ def pnorm(q, mean=0.0, sd=1.0):
72
+ q, mean, sd = _f(q), _f(mean), _f(sd)
73
+ with np.errstate(all="ignore"):
74
+ return 0.5 * erfc(-(q - mean) / (sd * math.sqrt(2.0)))
75
+
76
+
77
+ # Acklam's rational approximation to the normal quantile, |error| < 1.2e-9,
78
+ # then one Halley step against erfc for full double precision.
79
+ _QA = (-3.969683028665376e01, 2.209460984245205e02, -2.759285104469687e02,
80
+ 1.383577518672690e02, -3.066479806614716e01, 2.506628277459239e00)
81
+ _QB = (-5.447609879822406e01, 1.615858368580409e02, -1.556989798598866e02,
82
+ 6.680131188771972e01, -1.328068155288572e01)
83
+ _QC = (-7.784894002430293e-03, -3.223964580411365e-01, -2.400758277161838e00,
84
+ -2.549732539343734e00, 4.374664141464968e00, 2.938163982698783e00)
85
+ _QD = (7.784695709041462e-03, 3.224671290700398e-01, 2.445134137142996e00,
86
+ 3.754408661907416e00)
87
+
88
+
89
+ def _qnorm1(p: float) -> float:
90
+ if not 0.0 <= p <= 1.0 or math.isnan(p):
91
+ return math.nan
92
+ if p == 0.0:
93
+ return -math.inf
94
+ if p == 1.0:
95
+ return math.inf
96
+ low = 0.02425
97
+ if p < low:
98
+ q = math.sqrt(-2.0 * math.log(p))
99
+ x = (((((_QC[0] * q + _QC[1]) * q + _QC[2]) * q + _QC[3]) * q + _QC[4]) * q + _QC[5]) / (
100
+ (((_QD[0] * q + _QD[1]) * q + _QD[2]) * q + _QD[3]) * q + 1.0)
101
+ elif p <= 1.0 - low:
102
+ q = p - 0.5
103
+ r = q * q
104
+ x = (((((_QA[0] * r + _QA[1]) * r + _QA[2]) * r + _QA[3]) * r + _QA[4]) * r + _QA[5]) * q / (
105
+ ((((_QB[0] * r + _QB[1]) * r + _QB[2]) * r + _QB[3]) * r + _QB[4]) * r + 1.0)
106
+ else:
107
+ q = math.sqrt(-2.0 * math.log(1.0 - p))
108
+ x = -(((((_QC[0] * q + _QC[1]) * q + _QC[2]) * q + _QC[3]) * q + _QC[4]) * q + _QC[5]) / (
109
+ (((_QD[0] * q + _QD[1]) * q + _QD[2]) * q + _QD[3]) * q + 1.0)
110
+ err = 0.5 * math.erfc(-x / math.sqrt(2.0)) - p
111
+ u = err * math.sqrt(2.0 * math.pi) * math.exp(0.5 * x * x)
112
+ return x - u / (1.0 + 0.5 * x * u)
113
+
114
+
115
+ _qnorm_vec = _elementwise(_qnorm1)
116
+
117
+
118
+ def qnorm(p, mean=0.0, sd=1.0):
119
+ return _f(mean) + _f(sd) * _f(_qnorm_vec(p))
120
+
121
+
122
+ def dbeta(x, a, b):
123
+ x, a, b = _f(x), _f(a), _f(b)
124
+ with np.errstate(all="ignore"):
125
+ body = np.exp((a - 1.0) * np.log(x) + (b - 1.0) * np.log1p(-x) - _lbeta(a, b))
126
+ # x^0 at the edge is 1 even where log gives -inf * 0 = nan.
127
+ body = np.where((x == 0.0) & (a == 1.0), np.exp(-_lbeta(a, b)), body)
128
+ body = np.where((x == 1.0) & (b == 1.0), np.exp(-_lbeta(a, b)), body)
129
+ inside = (x >= 0.0) & (x <= 1.0)
130
+ return np.where(inside, body, 0.0)
131
+
132
+
133
+ def dt(x, df):
134
+ x, df = _f(x), _f(df)
135
+ with np.errstate(all="ignore"):
136
+ log_c = lgamma((df + 1.0) / 2.0) - lgamma(df / 2.0) - 0.5 * np.log(df * math.pi)
137
+ return np.exp(log_c - (df + 1.0) / 2.0 * np.log1p(x * x / df))
138
+
139
+
140
+ def dgamma(x, shape, rate=1.0):
141
+ x, shape, rate = _f(x), _f(shape), _f(rate)
142
+ with np.errstate(all="ignore"):
143
+ body = np.exp(
144
+ shape * np.log(rate) + (shape - 1.0) * np.log(x) - rate * x - lgamma(shape)
145
+ )
146
+ body = np.where((x == 0.0) & (shape == 1.0), rate, body)
147
+ return np.where(x >= 0.0, body, 0.0)
148
+
149
+
150
+ def dchisq(x, df):
151
+ return dgamma(x, _f(df) / 2.0, 0.5)
152
+
153
+
154
+ def dexp(x, rate=1.0):
155
+ x, rate = _f(x), _f(rate)
156
+ with np.errstate(all="ignore"):
157
+ return np.where(x >= 0.0, rate * np.exp(-rate * x), 0.0)
158
+
159
+
160
+ def dunif(x, min=0.0, max=1.0): # noqa: A002 - R's argument names
161
+ x, lo, hi = _f(x), _f(min), _f(max)
162
+ with np.errstate(all="ignore"):
163
+ return np.where((x >= lo) & (x <= hi), 1.0 / (hi - lo), 0.0)
164
+
165
+
166
+ def dlnorm(x, meanlog=0.0, sdlog=1.0):
167
+ x, mu, s = _f(x), _f(meanlog), _f(sdlog)
168
+ with np.errstate(all="ignore"):
169
+ z = (np.log(x) - mu) / s
170
+ body = np.exp(-0.5 * z * z) / (x * s * math.sqrt(2.0 * math.pi))
171
+ return np.where(x > 0.0, body, 0.0)
172
+
173
+
174
+ def _betacf(a: float, b: float, x: float) -> float:
175
+ """Continued fraction for the incomplete beta (modified Lentz)."""
176
+ tiny = 1e-300
177
+ qab, qap, qam = a + b, a + 1.0, a - 1.0
178
+ c, d = 1.0, 1.0 - qab * x / qap
179
+ d = 1.0 / (d if abs(d) > tiny else tiny)
180
+ h = d
181
+ for m in range(1, 400):
182
+ m2 = 2 * m
183
+ aa = m * (b - m) * x / ((qam + m2) * (a + m2))
184
+ d = 1.0 + aa * d
185
+ d = 1.0 / (d if abs(d) > tiny else tiny)
186
+ c = 1.0 + aa / c
187
+ c = c if abs(c) > tiny else tiny
188
+ h *= d * c
189
+ aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2))
190
+ d = 1.0 + aa * d
191
+ d = 1.0 / (d if abs(d) > tiny else tiny)
192
+ c = 1.0 + aa / c
193
+ c = c if abs(c) > tiny else tiny
194
+ delta = d * c
195
+ h *= delta
196
+ if abs(delta - 1.0) < 1e-15:
197
+ break
198
+ return h
199
+
200
+
201
+ def _betainc1(x: float, a: float, b: float) -> float:
202
+ """Regularized incomplete beta I_x(a, b)."""
203
+ if not (a > 0 and b > 0) or math.isnan(x):
204
+ return math.nan
205
+ if x <= 0.0:
206
+ return 0.0
207
+ if x >= 1.0:
208
+ return 1.0
209
+ log_front = (
210
+ math.lgamma(a + b) - math.lgamma(a) - math.lgamma(b)
211
+ + a * math.log(x) + b * math.log1p(-x)
212
+ )
213
+ front = math.exp(log_front)
214
+ if x < (a + 1.0) / (a + b + 2.0):
215
+ return front * _betacf(a, b, x) / a
216
+ return 1.0 - front * _betacf(b, a, 1.0 - x) / b
217
+
218
+
219
+ pbeta = _elementwise(_betainc1)
220
+
221
+
222
+ def _pt1(t: float, df: float) -> float:
223
+ if math.isnan(t) or not df > 0:
224
+ return math.nan
225
+ if math.isinf(t):
226
+ return 1.0 if t > 0 else 0.0
227
+ tail = 0.5 * _betainc1(df / (df + t * t), df / 2.0, 0.5)
228
+ return 1.0 - tail if t > 0 else tail
229
+
230
+
231
+ pt = _elementwise(_pt1)
232
+
233
+
234
+ def _qt1(p: float, df: float) -> float:
235
+ if math.isnan(p) or not df > 0 or not 0.0 <= p <= 1.0:
236
+ return math.nan
237
+ if p == 0.0:
238
+ return -math.inf
239
+ if p == 1.0:
240
+ return math.inf
241
+ if p == 0.5:
242
+ return 0.0
243
+ # Bracket, then bisect: robust for any df, and fast enough for a stat.
244
+ lo, hi = -1.0, 1.0
245
+ while _pt1(lo, df) > p:
246
+ lo *= 2.0
247
+ while _pt1(hi, df) < p:
248
+ hi *= 2.0
249
+ for _ in range(200):
250
+ mid = 0.5 * (lo + hi)
251
+ if _pt1(mid, df) < p:
252
+ lo = mid
253
+ else:
254
+ hi = mid
255
+ if hi - lo <= 1e-13 * max(1.0, abs(mid)):
256
+ break
257
+ return 0.5 * (lo + hi)
258
+
259
+
260
+ qt = _elementwise(_qt1)
261
+
262
+
263
+ def _qf1(p: float, df1: float, df2: float) -> float:
264
+ """Quantile of the F distribution, from the incomplete beta by bisection."""
265
+ if math.isnan(p) or not (df1 > 0 and df2 > 0) or not 0.0 <= p <= 1.0:
266
+ return math.nan
267
+ if p == 0.0:
268
+ return 0.0
269
+ if p == 1.0:
270
+ return math.inf
271
+ lo, hi = 0.0, 1.0
272
+ for _ in range(200):
273
+ mid = 0.5 * (lo + hi)
274
+ if _betainc1(mid, df1 / 2.0, df2 / 2.0) < p:
275
+ lo = mid
276
+ else:
277
+ hi = mid
278
+ if hi - lo <= 1e-15:
279
+ break
280
+ x = 0.5 * (lo + hi)
281
+ return df2 * x / (df1 * (1.0 - x))
282
+
283
+
284
+ qf = _elementwise(_qf1)
285
+
286
+
287
+ def pexp(q, rate=1.0):
288
+ q, rate = _f(q), _f(rate)
289
+ with np.errstate(all="ignore"):
290
+ return np.where(q > 0.0, -np.expm1(-rate * q), 0.0)
291
+
292
+
293
+ def punif(q, min=0.0, max=1.0): # noqa: A002 - R's argument names
294
+ q, lo, hi = _f(q), _f(min), _f(max)
295
+ with np.errstate(all="ignore"):
296
+ return np.clip((q - lo) / (hi - lo), 0.0, 1.0)
297
+
298
+
299
+ FUNCTIONS: dict[str, Callable[..., Any]] = {
300
+ "gamma": gamma,
301
+ "lgamma": lgamma,
302
+ "beta": beta,
303
+ "erf": erf,
304
+ "erfc": erfc,
305
+ "dnorm": dnorm,
306
+ "pnorm": pnorm,
307
+ "qnorm": qnorm,
308
+ "dbeta": dbeta,
309
+ "dt": dt,
310
+ "dgamma": dgamma,
311
+ "dchisq": dchisq,
312
+ "dexp": dexp,
313
+ "dunif": dunif,
314
+ "dlnorm": dlnorm,
315
+ "pbeta": pbeta,
316
+ "pt": pt,
317
+ "qt": qt,
318
+ "pexp": pexp,
319
+ "punif": punif,
320
+ }
321
+
322
+ # Densities integrate to 1, so a shaded area under one is a probability.
323
+ DENSITIES = frozenset({"dnorm", "dbeta", "dt", "dgamma", "dchisq", "dexp", "dunif", "dlnorm"})
324
+
325
+
326
+ def _support(name: str, args: list[float]) -> tuple[float, float] | None:
327
+ """The x range worth drawing for one density call (args after x)."""
328
+ def arg(i: int, default: float) -> float:
329
+ return float(args[i]) if len(args) > i else default
330
+
331
+ if name == "dnorm":
332
+ mu, sd = arg(0, 0.0), arg(1, 1.0)
333
+ return (mu - 4.0 * sd, mu + 4.0 * sd) if sd > 0 else None
334
+ if name == "dbeta":
335
+ return (0.0, 1.0)
336
+ if name == "dt":
337
+ df = arg(0, 1.0)
338
+ half = 5.0 if df >= 3 else 8.0
339
+ return (-half, half)
340
+ if name in {"dgamma", "dchisq"}:
341
+ if name == "dgamma":
342
+ shape, rate = arg(0, 1.0), arg(1, 1.0)
343
+ else:
344
+ shape, rate = arg(0, 1.0) / 2.0, 0.5
345
+ if shape <= 0 or rate <= 0:
346
+ return None
347
+ mean, sd = shape / rate, math.sqrt(shape) / rate
348
+ return (0.0, mean + 5.0 * sd)
349
+ if name == "dexp":
350
+ rate = arg(0, 1.0)
351
+ return (0.0, 6.0 / rate) if rate > 0 else None
352
+ if name == "dunif":
353
+ lo, hi = arg(0, 0.0), arg(1, 1.0)
354
+ pad = 0.1 * (hi - lo)
355
+ return (lo - pad, hi + pad) if hi > lo else None
356
+ if name == "dlnorm":
357
+ mu, s = arg(0, 0.0), arg(1, 1.0)
358
+ return (0.0, math.exp(mu + 3.0 * s)) if s > 0 else None
359
+ return None
360
+
361
+
362
+ def density_calls(tree: ast.AST | None, variable: str) -> list[ast.Call]:
363
+ """Density calls whose first argument is the plot variable itself."""
364
+ if tree is None:
365
+ return []
366
+ return [
367
+ node
368
+ for node in ast.walk(tree)
369
+ if isinstance(node, ast.Call)
370
+ and isinstance(node.func, ast.Name)
371
+ and node.func.id in DENSITIES
372
+ and node.args
373
+ and isinstance(node.args[0], ast.Name)
374
+ and node.args[0].id == variable
375
+ ]
376
+
377
+
378
+ def support_hint(
379
+ tree: ast.AST | None, variable: str, namespace: dict[str, Any]
380
+ ) -> tuple[float, float] | None:
381
+ """Union of the supports of every density of ``variable`` in ``tree``.
382
+
383
+ None when there is no density, or a parameter is not a plain number
384
+ (a slider or transition changes it frame by frame).
385
+ """
386
+ calls = density_calls(tree, variable)
387
+ if not calls:
388
+ return None
389
+ lo, hi = math.inf, -math.inf
390
+ for call in calls:
391
+ values = []
392
+ for node in call.args[1:]:
393
+ try:
394
+ code = compile(ast.Expression(body=node), "<support>", "eval")
395
+ value = float(eval(code, {"__builtins__": {}}, dict(namespace))) # noqa: S307
396
+ except Exception:
397
+ return None
398
+ if not math.isfinite(value):
399
+ return None
400
+ values.append(value)
401
+ span = _support(call.func.id, values)
402
+ if span is None:
403
+ return None
404
+ lo, hi = min(lo, span[0]), max(hi, span[1])
405
+ if not (math.isfinite(lo) and math.isfinite(hi)) or hi <= lo:
406
+ return None
407
+ return lo, hi