1use crate::error::{OpError, OpResult};
8use std::f64::consts::{PI, SQRT_2};
9use std::fmt;
10
11pub const Z95: f64 = 1.644853626951472;
14
15#[derive(Clone, Copy, Debug, PartialEq)]
17pub enum Family {
18 Normal {
19 mean: f64,
20 sd: f64,
21 },
22 Lognormal {
23 mu: f64,
24 sigma: f64,
25 },
26 Uniform {
27 lo: f64,
28 hi: f64,
29 },
30 Beta {
31 a: f64,
32 b: f64,
33 },
34 Gamma {
35 shape: f64,
36 scale: f64,
37 },
38 Exponential {
39 rate: f64,
40 },
41 Triangular {
42 lo: f64,
43 mode: f64,
44 hi: f64,
45 },
46 Pert {
48 lo: f64,
49 mode: f64,
50 hi: f64,
51 },
52}
53
54fn check(ok: bool, message: impl FnOnce() -> String) -> OpResult<()> {
55 if ok { Ok(()) } else { Err(OpError::new(message())) }
56}
57
58fn finite(values: &[f64], what: &str) -> OpResult<()> {
59 check(values.iter().all(|x| x.is_finite()), || {
60 format!("{what} needs finite numbers")
61 })
62}
63
64fn ratio(a: f64, b: f64) -> f64 {
66 if a >= b {
67 1.0 / (1.0 + b / a)
68 } else {
69 let r = a / b;
70 r / (1.0 + r)
71 }
72}
73
74fn interval_fraction(lo: f64, hi: f64, x: f64) -> f64 {
75 if x <= lo {
76 return 0.0;
77 }
78 if x >= hi {
79 return 1.0;
80 }
81 let width = hi - lo;
82 if width.is_finite() {
83 (x - lo) / width
84 } else {
85 (x / 2.0 - lo / 2.0) / (hi / 2.0 - lo / 2.0)
86 }
87}
88
89fn inverse_width(lo: f64, hi: f64) -> f64 {
90 let width = hi - lo;
91 if width.is_finite() {
92 1.0 / width
93 } else {
94 0.5 / (hi / 2.0 - lo / 2.0)
95 }
96}
97
98fn standardized(x: f64, mean: f64, sd: f64) -> f64 {
99 if !x.is_finite() {
100 return x;
101 }
102 let delta = x - mean;
103 if delta.is_finite() {
104 delta / sd
105 } else {
106 x / sd - mean / sd
107 }
108}
109
110impl Family {
111 pub fn normal(mean: f64, sd: f64) -> OpResult<Family> {
112 finite(&[mean, sd], "normal")?;
113 check(sd > 0.0, || "normal needs a standard deviation above 0".into())?;
114 Ok(Family::Normal { mean, sd })
115 }
116
117 pub fn lognormal(mu: f64, sigma: f64) -> OpResult<Family> {
118 finite(&[mu, sigma], "lognormal")?;
119 check(sigma > 0.0, || "lognormal needs a sigma above 0".into())?;
120 Ok(Family::Lognormal { mu, sigma })
121 }
122
123 pub fn uniform(lo: f64, hi: f64) -> OpResult<Family> {
124 finite(&[lo, hi], "uniform")?;
125 check(lo < hi, || "uniform needs its lower end below its upper end".into())?;
126 Ok(Family::Uniform { lo, hi })
127 }
128
129 pub fn beta(a: f64, b: f64) -> OpResult<Family> {
130 finite(&[a, b], "beta")?;
131 check(a > 0.0 && b > 0.0, || "beta needs two parameters above 0".into())?;
132 Ok(Family::Beta { a, b })
133 }
134
135 pub fn gamma(shape: f64, scale: f64) -> OpResult<Family> {
136 finite(&[shape, scale], "gamma")?;
137 check(shape > 0.0 && scale > 0.0, || {
138 "gamma needs a shape and a scale above 0".into()
139 })?;
140 Ok(Family::Gamma { shape, scale })
141 }
142
143 pub fn exponential(rate: f64) -> OpResult<Family> {
144 finite(&[rate], "exponential")?;
145 check(rate > 0.0, || "exponential needs a rate above 0".into())?;
146 Ok(Family::Exponential { rate })
147 }
148
149 pub fn triangular(lo: f64, mode: f64, hi: f64) -> OpResult<Family> {
150 finite(&[lo, mode, hi], "triangular")?;
151 check(lo < hi && lo <= mode && mode <= hi, || {
152 "triangular needs lo < hi, with the mode between them".into()
153 })?;
154 Ok(Family::Triangular { lo, mode, hi })
155 }
156
157 pub fn pert(lo: f64, mode: f64, hi: f64) -> OpResult<Family> {
158 finite(&[lo, mode, hi], "pert")?;
159 check(lo < hi && lo <= mode && mode <= hi, || {
160 "pert needs lo < hi, with the mode between them".into()
161 })?;
162 Ok(Family::Pert { lo, mode, hi })
163 }
164
165 pub fn estimate(a: f64, b: f64) -> OpResult<Family> {
167 finite(&[a, b], "`to`")?;
168 if a <= 0.0 || b <= 0.0 {
169 return Err(OpError::new("`a to b` needs two positive numbers")
170 .help("for a quantity that can be zero or negative, use `normal_range(lo, hi)`"));
171 }
172 check(a < b, || "`a to b` needs a below b".into())?;
173 let (la, lb) = (libm::log(a), libm::log(b));
174 Family::lognormal((la + lb) / 2.0, (lb - la) / (2.0 * Z95))
175 }
176
177 pub fn normal_range(lo: f64, hi: f64) -> OpResult<Family> {
180 finite(&[lo, hi], "normal_range")?;
181 check(lo < hi, || {
182 "normal_range needs its lower end below its upper end".into()
183 })?;
184 Family::normal(
185 crate::stats::midpoint(lo, hi),
186 crate::stats::scaled_difference(hi, lo, 1.0 / (2.0 * Z95)),
187 )
188 }
189
190 fn pert_shape(lo: f64, mode: f64, hi: f64) -> (f64, f64) {
192 let p = interval_fraction(lo, hi, mode);
193 (1.0 + 4.0 * p, 1.0 + 4.0 * (1.0 - p))
194 }
195
196 pub fn support(&self) -> (f64, f64) {
198 match *self {
199 Family::Normal { .. } => (f64::NEG_INFINITY, f64::INFINITY),
200 Family::Lognormal { .. } | Family::Gamma { .. } | Family::Exponential { .. } => (0.0, f64::INFINITY),
201 Family::Uniform { lo, hi } | Family::Triangular { lo, hi, .. } | Family::Pert { lo, hi, .. } => (lo, hi),
202 Family::Beta { .. } => (0.0, 1.0),
203 }
204 }
205
206 pub fn mean(&self) -> f64 {
207 match *self {
208 Family::Normal { mean, .. } => mean,
209 Family::Lognormal { mu, sigma } => crate::math::exp(mu + sigma * sigma / 2.0),
210 Family::Uniform { lo, hi } => crate::stats::midpoint(lo, hi),
211 Family::Beta { a, b } => ratio(a, b),
212 Family::Gamma { shape, scale } => shape * scale,
213 Family::Exponential { rate } => 1.0 / rate,
214 Family::Triangular { lo, mode, hi } => crate::stats::lerp(crate::stats::midpoint(lo, hi), mode, 1.0 / 3.0),
215 Family::Pert { lo, mode, hi } => {
216 let (a, b) = Family::pert_shape(lo, mode, hi);
217 crate::stats::lerp(lo, hi, ratio(a, b))
218 }
219 }
220 }
221
222 pub fn variance(&self) -> f64 {
223 let sd = self.sd();
224 sd * sd
225 }
226
227 pub fn sd(&self) -> f64 {
228 match *self {
229 Family::Normal { sd, .. } => sd,
230 Family::Lognormal { mu, sigma } => {
231 let s2 = sigma * sigma;
232 let log_sd = if s2 == 0.0 {
233 mu + libm::log(sigma)
234 } else {
235 mu + s2 + 0.5 * libm::log(-libm::expm1(-s2))
236 };
237 crate::math::exp(log_sd)
238 }
239 Family::Uniform { lo, hi } => crate::stats::scaled_difference(hi, lo, 1.0 / libm::sqrt(12.0)),
240 Family::Beta { a, b } => {
241 let max = a.max(b);
242 let inv = if max < 1.0 {
243 1.0 / (1.0 + a + b)
244 } else {
245 (1.0 / max) / (1.0 + a.min(b) / max + 1.0 / max)
246 };
247 libm::sqrt(ratio(a, b)) * libm::sqrt(ratio(b, a)) * libm::sqrt(inv)
248 }
249 Family::Gamma { shape, scale } => libm::sqrt(shape) * scale,
250 Family::Exponential { rate } => 1.0 / rate,
251 Family::Triangular { lo, mode, hi } => {
252 let d = |a, b| crate::stats::scaled_difference(a, b, 1.0 / 6.0);
253 libm::hypot(libm::hypot(d(mode, lo), d(hi, mode)), d(hi, lo))
254 }
255 Family::Pert { lo, mode, hi } => {
256 let (a, b) = Family::pert_shape(lo, mode, hi);
257 crate::stats::scaled_difference(hi, lo, Family::Beta { a, b }.sd())
258 }
259 }
260 }
261
262 pub fn interval_moments(&self, lo: f64, hi: f64) -> (f64, f64) {
266 let mass = self.cdf(hi) - self.cdf(lo);
267 let raw = |first: f64, second: f64| {
268 let mean = first / mass;
269 (mean, (second / mass - mean * mean).max(0.0))
270 };
271 match *self {
272 Family::Uniform { .. } => (crate::stats::midpoint(lo, hi), Family::Uniform { lo, hi }.variance()),
273 Family::Normal { mean, sd } => {
274 let (a, b) = ((lo - mean) / sd, (hi - mean) / sd);
275 let (pa, pb) = (std_normal_pdf(a), std_normal_pdf(b));
276 let shift = (pa - pb) / mass;
277 let edge = |z: f64, p: f64| if z.is_finite() { z * p } else { 0.0 };
278 (
279 mean + sd * shift,
280 sd * sd * (1.0 + (edge(a, pa) - edge(b, pb)) / mass - shift * shift).max(0.0),
281 )
282 }
283 Family::Beta { a, b } => {
284 let moment = |n: f64| beta_cdf(a + n, b, hi) - beta_cdf(a + n, b, lo);
285 raw(
286 a / (a + b) * moment(1.0),
287 a * (a + 1.0) / ((a + b) * (a + b + 1.0)) * moment(2.0),
288 )
289 }
290 Family::Gamma { shape, scale } => {
291 let moment = |n: f64| gamma_cdf(shape + n, hi / scale) - gamma_cdf(shape + n, lo / scale);
292 raw(
293 shape * scale * moment(1.0),
294 shape * (shape + 1.0) * scale * scale * moment(2.0),
295 )
296 }
297 Family::Exponential { rate } => Family::Gamma {
298 shape: 1.0,
299 scale: 1.0 / rate,
300 }
301 .interval_moments(lo, hi),
302 Family::Lognormal { mu, sigma } => {
303 let moment = |n: f64| {
304 let cdf = |x| std_normal_cdf((libm::log(x) - mu - n * sigma * sigma) / sigma);
305 crate::math::exp(n * mu + n * n * sigma * sigma / 2.0) * (cdf(hi) - cdf(lo))
306 };
307 raw(moment(1.0), moment(2.0))
308 }
309 Family::Pert { lo: a, mode, hi: b } => {
310 let (alpha, beta) = Self::pert_shape(a, mode, b);
311 let (m, v) =
312 Family::Beta { a: alpha, b: beta }.interval_moments((lo - a) / (b - a), (hi - a) / (b - a));
313 (a + (b - a) * m, (b - a).powi(2) * v)
314 }
315 Family::Triangular { lo: a, mode, hi: b } => {
316 let width = b - a;
317 let (l, h, m) = ((lo - a) / width, (hi - a) / width, (mode - a) / width);
318 let moment = |n: i32| {
319 let integral = |l: f64, h: f64, k: i32| (h.powi(k + 1) - l.powi(k + 1)) / (k + 1) as f64;
320 let left = if l < m {
321 2.0 / m * integral(l, h.min(m), n + 1)
322 } else {
323 0.0
324 };
325 let right = if h > m {
326 2.0 / (1.0 - m) * (integral(l.max(m), h, n) - integral(l.max(m), h, n + 1))
327 } else {
328 0.0
329 };
330 left + right
331 };
332 let (m, v) = raw(moment(1), moment(2));
333 (a + width * m, width * width * v)
334 }
335 }
336 }
337
338 pub fn pdf(&self, x: f64) -> f64 {
339 match *self {
340 Family::Normal { mean, sd } => std_normal_pdf(standardized(x, mean, sd)) / sd,
341 Family::Lognormal { mu, sigma } => {
342 if x <= 0.0 {
343 0.0
344 } else {
345 std_normal_pdf((libm::log(x) - mu) / sigma) / (sigma * x)
346 }
347 }
348 Family::Uniform { lo, hi } => {
349 if (lo..=hi).contains(&x) {
350 inverse_width(lo, hi)
351 } else {
352 0.0
353 }
354 }
355 Family::Beta { a, b } => beta_pdf(a, b, x),
356 Family::Gamma { shape, scale } => {
357 if x < 0.0 {
358 return 0.0;
359 }
360 if x == 0.0 {
361 return if shape < 1.0 {
362 f64::INFINITY
363 } else if shape == 1.0 {
364 1.0 / scale
365 } else {
366 0.0
367 };
368 }
369 let y = x / scale;
370 crate::math::exp((shape - 1.0) * libm::log(y) - y - libm::lgamma(shape)) / scale
371 }
372 Family::Exponential { rate } => {
373 if x < 0.0 {
374 0.0
375 } else {
376 rate * crate::math::exp(-rate * x)
377 }
378 }
379 Family::Triangular { lo, mode, hi } => {
380 if x < lo || x > hi {
381 0.0
382 } else if x < mode {
383 2.0 * interval_fraction(lo, mode, x) * inverse_width(lo, hi)
384 } else if x > mode {
385 2.0 * (1.0 - interval_fraction(mode, hi, x)) * inverse_width(lo, hi)
386 } else {
387 2.0 * inverse_width(lo, hi)
388 }
389 }
390 Family::Pert { lo, mode, hi } => {
391 let (a, b) = Family::pert_shape(lo, mode, hi);
392 if x < lo || x > hi {
393 0.0
394 } else {
395 beta_pdf(a, b, interval_fraction(lo, hi, x)) * inverse_width(lo, hi)
396 }
397 }
398 }
399 }
400
401 pub fn cdf(&self, x: f64) -> f64 {
403 if x.is_nan() {
404 return f64::NAN;
405 }
406 match *self {
407 Family::Normal { mean, sd } => std_normal_cdf(standardized(x, mean, sd)),
408 Family::Lognormal { mu, sigma } => {
409 if x <= 0.0 {
410 0.0
411 } else {
412 std_normal_cdf((libm::log(x) - mu) / sigma)
413 }
414 }
415 Family::Uniform { lo, hi } => interval_fraction(lo, hi, x).clamp(0.0, 1.0),
416 Family::Beta { a, b } => beta_cdf(a, b, x),
417 Family::Gamma { shape, scale } => gamma_cdf(shape, x / scale),
418 Family::Exponential { rate } => {
419 if x <= 0.0 {
420 0.0
421 } else {
422 -libm::expm1(-rate * x)
423 }
424 }
425 Family::Triangular { lo, mode, hi } => {
426 if x <= lo {
427 0.0
428 } else if x >= hi {
429 1.0
430 } else if x <= mode {
431 interval_fraction(lo, hi, x) * interval_fraction(lo, mode, x)
432 } else {
433 1.0 - (1.0 - interval_fraction(lo, hi, x)) * (1.0 - interval_fraction(mode, hi, x))
434 }
435 }
436 Family::Pert { lo, mode, hi } => {
437 let (a, b) = Family::pert_shape(lo, mode, hi);
438 beta_cdf(a, b, interval_fraction(lo, hi, x))
439 }
440 }
441 }
442
443 pub fn quantile(&self, p: f64) -> f64 {
445 let (lo, hi) = self.support();
446 if p <= 0.0 {
447 return lo;
448 }
449 if p >= 1.0 {
450 return hi;
451 }
452 match *self {
453 Family::Normal { mean, sd } => mean + sd * std_normal_quantile(p),
454 Family::Lognormal { mu, sigma } => crate::math::exp(mu + sigma * std_normal_quantile(p)),
455 Family::Uniform { lo, hi } => crate::stats::lerp(lo, hi, p),
456 Family::Exponential { rate } => -libm::log1p(-p) / rate,
457 Family::Triangular { lo, mode, hi } => {
458 let split = interval_fraction(lo, hi, mode);
459 if p <= split {
460 crate::stats::lerp(lo, hi, libm::sqrt(p * split))
461 } else {
462 crate::stats::lerp(hi, lo, libm::sqrt((1.0 - p) * (1.0 - split)))
463 }
464 }
465 Family::Beta { .. } | Family::Pert { .. } => invert(|x| self.cdf(x), p, lo, hi),
466 Family::Gamma { .. } => {
467 let mut top = self.mean() + 10.0 * libm::sqrt(self.variance());
468 while self.cdf(top) < p && top < f64::MAX / 4.0 {
469 top *= 2.0;
470 }
471 invert(|x| self.cdf(x), p, 0.0, top)
472 }
473 }
474 }
475
476 pub fn sample(&self, rng: &mut Rng) -> f64 {
478 match *self {
479 Family::Normal { mean, sd } => mean + sd * rng.normal(),
480 Family::Lognormal { mu, sigma } => crate::math::exp(mu + sigma * rng.normal()),
481 Family::Uniform { lo, hi } => crate::stats::lerp(lo, hi, rng.uniform()),
482 Family::Beta { a, b } => rng.beta(a, b),
483 Family::Gamma { shape, scale } => scale * rng.gamma(shape),
484 Family::Exponential { rate } => -libm::log(rng.open()) / rate,
485 Family::Triangular { .. } => self.quantile(rng.uniform()),
486 Family::Pert { lo, mode, hi } => {
487 let (a, b) = Family::pert_shape(lo, mode, hi);
488 crate::stats::lerp(lo, hi, rng.beta(a, b))
489 }
490 }
491 }
492
493 pub fn name(&self) -> &'static str {
494 match self {
495 Family::Normal { .. } => "normal",
496 Family::Lognormal { .. } => "lognormal",
497 Family::Uniform { .. } => "uniform",
498 Family::Beta { .. } => "beta",
499 Family::Gamma { .. } => "gamma",
500 Family::Exponential { .. } => "exponential",
501 Family::Triangular { .. } => "triangular",
502 Family::Pert { .. } => "pert",
503 }
504 }
505
506 pub fn params(&self) -> Vec<f64> {
508 match *self {
509 Family::Normal { mean, sd } => vec![mean, sd],
510 Family::Lognormal { mu, sigma } => vec![mu, sigma],
511 Family::Uniform { lo, hi } => vec![lo, hi],
512 Family::Beta { a, b } => vec![a, b],
513 Family::Gamma { shape, scale } => vec![shape, scale],
514 Family::Exponential { rate } => vec![rate],
515 Family::Triangular { lo, mode, hi } | Family::Pert { lo, mode, hi } => vec![lo, mode, hi],
516 }
517 }
518}
519
520impl fmt::Display for Family {
521 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
522 let params: Vec<String> = self
523 .params()
524 .iter()
525 .map(|x| {
526 format!("{:.4}", x)
527 .trim_end_matches('0')
528 .trim_end_matches('.')
529 .to_string()
530 })
531 .collect();
532 write!(f, "{}({})", self.name(), params.join(", "))
533 }
534}
535
536fn std_normal_pdf(z: f64) -> f64 {
539 crate::math::exp(-z * z / 2.0) / libm::sqrt(2.0 * PI)
540}
541
542pub fn std_normal_cdf(z: f64) -> f64 {
543 0.5 * libm::erfc(-z / SQRT_2)
544}
545
546pub fn std_normal_quantile(p: f64) -> f64 {
549 if p == 0.5 {
550 return 0.0;
551 }
552 if p <= 0.0 {
553 return f64::NEG_INFINITY;
554 }
555 if p >= 1.0 {
556 return f64::INFINITY;
557 }
558 let tail = |q: f64| {
559 let t = libm::sqrt(-2.0 * libm::log(q));
560 t - (2.515517 + 0.802853 * t + 0.010328 * t * t)
561 / (1.0 + 1.432788 * t + 0.189269 * t * t + 0.001308 * t * t * t)
562 };
563 let mut z = if p < 0.5 { -tail(p) } else { tail(1.0 - p) };
564 for _ in 0..6 {
565 let density = std_normal_pdf(z);
566 if density == 0.0 {
567 break;
568 }
569 let step = (std_normal_cdf(z) - p) / density;
570 z -= step;
571 if step.abs() < 1e-15 * z.abs().max(1.0) {
572 break;
573 }
574 }
575 z
576}
577
578fn beta_pdf(a: f64, b: f64, x: f64) -> f64 {
579 if !(0.0..=1.0).contains(&x) {
580 return 0.0;
581 }
582 if x == 0.0 || x == 1.0 {
583 let edge = if x == 0.0 { a } else { b };
586 return if edge < 1.0 {
587 f64::INFINITY
588 } else if edge == 1.0 {
589 crate::math::exp(libm::lgamma(a + b) - libm::lgamma(a) - libm::lgamma(b))
590 } else {
591 0.0
592 };
593 }
594 crate::math::exp(
595 (a - 1.0) * libm::log(x) + (b - 1.0) * libm::log1p(-x) + libm::lgamma(a + b)
596 - libm::lgamma(a)
597 - libm::lgamma(b),
598 )
599}
600
601const LN_SQRT_2PI: f64 = 0.918_938_533_204_672_7;
603
604pub fn ln_beta(a: f64, b: f64) -> f64 {
608 let (p, q) = if a < b { (a, b) } else { (b, a) };
609 let share = p / (p + q);
610 if p >= 10.0 {
611 let rest = stirling_rest(p) + stirling_rest(q) - stirling_rest(p + q);
612 -0.5 * libm::log(q) + LN_SQRT_2PI + rest + (p - 0.5) * libm::log(share) + q * libm::log1p(-share)
613 } else if q >= 10.0 {
614 let rest = stirling_rest(q) - stirling_rest(p + q);
615 libm::lgamma(p) + rest + p - p * libm::log(p + q) + (q - 0.5) * libm::log1p(-share)
616 } else {
617 libm::lgamma(p) + libm::lgamma(q) - libm::lgamma(p + q)
618 }
619}
620
621fn stirling_rest(x: f64) -> f64 {
625 let r = 1.0 / (x * x);
626 let series = 1.0 / 12.0
627 + r * (-1.0 / 360.0
628 + r * (1.0 / 1260.0 + r * (-1.0 / 1680.0 + r * (1.0 / 1188.0 + r * (-691.0 / 360_360.0 + r / 156.0)))));
629 series / x
630}
631
632pub fn beta_cdf(a: f64, b: f64, x: f64) -> f64 {
635 if x <= 0.0 {
636 return 0.0;
637 }
638 if x >= 1.0 {
639 return 1.0;
640 }
641 let front = crate::math::exp(
642 libm::lgamma(a + b) - libm::lgamma(a) - libm::lgamma(b) + a * libm::log(x) + b * libm::log1p(-x),
643 );
644 if x < (a + 1.0) / (a + b + 2.0) {
645 front * beta_fraction(a, b, x) / a
646 } else {
647 1.0 - front * beta_fraction(b, a, 1.0 - x) / b
648 }
649}
650
651fn beta_fraction(a: f64, b: f64, x: f64) -> f64 {
652 const TINY: f64 = 1e-300;
653 let (qab, qap, qam) = (a + b, a + 1.0, a - 1.0);
654 let mut c = 1.0;
655 let mut d = 1.0 - qab * x / qap;
656 if d.abs() < TINY {
657 d = TINY;
658 }
659 d = 1.0 / d;
660 let mut h = d;
661 for m in 1..=1000 {
662 let m = m as f64;
663 let m2 = 2.0 * m;
664 let aa = m * (b - m) * x / ((qam + m2) * (a + m2));
665 d = 1.0 + aa * d;
666 if d.abs() < TINY {
667 d = TINY;
668 }
669 c = 1.0 + aa / c;
670 if c.abs() < TINY {
671 c = TINY;
672 }
673 d = 1.0 / d;
674 h *= d * c;
675 let aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2));
676 d = 1.0 + aa * d;
677 if d.abs() < TINY {
678 d = TINY;
679 }
680 c = 1.0 + aa / c;
681 if c.abs() < TINY {
682 c = TINY;
683 }
684 d = 1.0 / d;
685 let delta = d * c;
686 h *= delta;
687 if (delta - 1.0).abs() < 1e-16 {
688 break;
689 }
690 }
691 h
692}
693
694pub fn gamma_cdf(a: f64, x: f64) -> f64 {
697 if x == f64::INFINITY {
698 return 1.0;
699 }
700 if x <= 0.0 {
701 return 0.0;
702 }
703 let front = crate::math::exp(-x + a * libm::log(x) - libm::lgamma(a));
704 if x < a + 1.0 {
705 let (mut sum, mut term, mut n) = (1.0 / a, 1.0 / a, a);
706 for _ in 0..10_000 {
707 n += 1.0;
708 term *= x / n;
709 sum += term;
710 if term.abs() < sum.abs() * 1e-17 {
711 break;
712 }
713 }
714 (sum * front).min(1.0)
715 } else {
716 const TINY: f64 = 1e-300;
717 let mut b = x + 1.0 - a;
718 let mut c = 1.0 / TINY;
719 let mut d = 1.0 / b;
720 let mut h = d;
721 for i in 1..10_000 {
722 let i = i as f64;
723 let an = -i * (i - a);
724 b += 2.0;
725 d = an * d + b;
726 if d.abs() < TINY {
727 d = TINY;
728 }
729 c = b + an / c;
730 if c.abs() < TINY {
731 c = TINY;
732 }
733 d = 1.0 / d;
734 let delta = d * c;
735 h *= delta;
736 if (delta - 1.0).abs() < 1e-16 {
737 break;
738 }
739 }
740 (1.0 - front * h).max(0.0)
741 }
742}
743
744pub fn invert(f: impl Fn(f64) -> f64, p: f64, mut lo: f64, mut hi: f64) -> f64 {
747 if !lo.is_finite() || !hi.is_finite() {
748 return f64::NAN;
749 }
750 for _ in 0..300 {
751 let mid = crate::stats::midpoint(lo, hi);
752 if mid <= lo || mid >= hi {
753 break;
754 }
755 let cumulative = f(mid);
756 if !cumulative.is_finite() {
757 return f64::NAN;
758 }
759 if cumulative >= p {
760 hi = mid;
761 } else {
762 lo = mid;
763 }
764 }
765 hi
766}
767
768#[derive(Clone, Debug)]
772pub struct Rng {
773 s: [u64; 4],
774}
775
776const GOLDEN: u64 = 0x9E37_79B9_7F4A_7C15;
778
779pub(crate) fn mix(z: u64) -> u64 {
780 let z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
781 let z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
782 z ^ (z >> 31)
783}
784
785impl Rng {
786 pub fn new(seed: u64) -> Rng {
787 let mut z = seed;
788 let mut next = || {
789 z = z.wrapping_add(GOLDEN);
790 mix(z)
791 };
792 Rng {
793 s: [next(), next(), next(), next()],
794 }
795 }
796
797 pub fn stream(seed: u64, index: u64) -> Rng {
802 Rng::new(mix(seed.wrapping_add(index.wrapping_add(1).wrapping_mul(GOLDEN))))
805 }
806
807 pub fn next_u64(&mut self) -> u64 {
808 let s = &mut self.s;
809 let result = s[0].wrapping_add(s[3]).rotate_left(23).wrapping_add(s[0]);
810 let t = s[1] << 17;
811 s[2] ^= s[0];
812 s[3] ^= s[1];
813 s[1] ^= s[2];
814 s[0] ^= s[3];
815 s[2] ^= t;
816 s[3] = s[3].rotate_left(45);
817 result
818 }
819
820 pub fn uniform(&mut self) -> f64 {
822 (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
823 }
824
825 pub fn open(&mut self) -> f64 {
827 1.0 - self.uniform()
828 }
829
830 pub fn normal(&mut self) -> f64 {
832 let (u, v) = (self.open(), self.uniform());
833 libm::sqrt(-2.0 * libm::log(u)) * libm::cos(2.0 * PI * v)
834 }
835
836 pub fn gamma(&mut self, shape: f64) -> f64 {
838 if shape < 1.0 {
839 let u = self.open();
840 return self.gamma(shape + 1.0) * libm::pow(u, 1.0 / shape);
841 }
842 let d = shape - 1.0 / 3.0;
843 let c = 1.0 / libm::sqrt(9.0 * d);
844 loop {
845 let x = self.normal();
846 let v = 1.0 + c * x;
847 if v <= 0.0 {
848 continue;
849 }
850 let v = v * v * v;
851 let u = self.open();
852 if u < 1.0 - 0.0331 * x.powi(4) || libm::log(u) < 0.5 * x * x + d * (1.0 - v + libm::log(v)) {
853 return d * v;
854 }
855 }
856 }
857
858 pub fn beta(&mut self, a: f64, b: f64) -> f64 {
859 for _ in 0..16 {
860 let x = self.gamma(a);
861 let y = self.gamma(b);
862 if x + y > 0.0 {
863 return if x == 0.0 {
864 0.0
865 } else if y == 0.0 {
866 1.0
867 } else {
868 ratio(x, y)
869 };
870 }
871 }
872 if self.uniform() < ratio(a, b) { 1.0 } else { 0.0 }
875 }
876
877 pub fn choose(&mut self, weights: impl Iterator<Item = f64> + Clone) -> Option<usize> {
880 let total: f64 = weights.clone().sum();
881 if total <= 0.0 {
882 return None;
883 }
884 let target = self.uniform() * total;
885 let mut acc = 0.0;
886 let mut last = None;
887 for (i, w) in weights.enumerate() {
888 if w <= 0.0 {
889 continue;
890 }
891 acc += w;
892 last = Some(i);
893 if target < acc {
894 return Some(i);
895 }
896 }
897 last
898 }
899}
900
901#[derive(Clone, Debug)]
905pub enum Part {
906 Point(f64),
907 Continuous(Family),
908 Analytic(crate::analytic::Analytic),
909}
910
911#[derive(Clone, Debug)]
914pub struct Mixture {
915 pub parts: Vec<(Part, f64)>,
916}
917
918impl Mixture {
919 fn total(&self) -> f64 {
920 crate::stats::sum(self.parts.iter().map(|(_, p)| *p))
921 }
922
923 pub fn mean(&self) -> f64 {
924 crate::stats::weighted_mean(self.parts.iter().map(|(part, p)| {
925 (
926 match part {
927 Part::Point(x) => *x,
928 Part::Continuous(f) => f.mean(),
929 Part::Analytic(a) => a.moments().0,
930 },
931 *p,
932 )
933 }))
934 }
935
936 pub fn variance(&self) -> f64 {
937 let sd = self.sd();
938 sd * sd
939 }
940
941 pub fn sd(&self) -> f64 {
942 crate::stats::weighted_sd(self.parts.iter().map(|(part, p)| {
943 let (m, sd) = match part {
944 Part::Point(x) => (*x, 0.0),
945 Part::Continuous(f) => (f.mean(), f.sd()),
946 Part::Analytic(a) => (a.moments().0, a.sd()),
947 };
948 (m, sd, *p)
949 }))
950 }
951
952 pub fn cdf(&self, x: f64) -> f64 {
953 let below = crate::stats::sum(self.parts.iter().map(|(part, p)| {
954 p * match part {
955 Part::Point(v) => {
956 if *v <= x {
957 1.0
958 } else {
959 0.0
960 }
961 }
962 Part::Continuous(f) => f.cdf(x),
963 Part::Analytic(a) => a.cdf(x),
964 }
965 }));
966 below / self.total()
967 }
968
969 pub fn median_bounds(&self) -> (f64, f64) {
973 let mut intervals = Vec::new();
974 for (part, weight) in &self.parts {
975 if *weight <= 0.0 {
976 continue;
977 }
978 match part {
979 Part::Point(x) => intervals.push((*x, *x, *weight)),
980 Part::Continuous(f) => {
981 let (lo, hi) = f.support();
982 intervals.push((lo, hi, *weight));
983 }
984 Part::Analytic(a) => {
985 let total = a.domain.mass();
986 for &(lo, hi) in &a.domain.0 {
987 let x = a.scale * a.family.quantile(lo) + a.offset;
988 let y = a.scale * a.family.quantile(hi) + a.offset;
989 intervals.push((x.min(y), x.max(y), weight * (hi - lo) / total));
990 }
991 }
992 }
993 }
994 intervals.sort_by(|a, b| a.0.total_cmp(&b.0));
995 let total = crate::stats::sum(intervals.iter().map(|x| x.2));
996 let mut acc = crate::stats::Sum::default();
997 let mut end = f64::NEG_INFINITY;
998 for (lo, hi, weight) in intervals {
999 if lo > end && acc.value() > 0.0 && crate::stats::half_split(acc.value(), total) {
1000 return (end, lo);
1001 }
1002 end = end.max(hi);
1003 acc.add(weight);
1004 }
1005 let x = self.quantile(0.5);
1006 (x, x)
1007 }
1008
1009 pub fn median(&self) -> f64 {
1010 let (lo, hi) = self.median_bounds();
1011 crate::stats::midpoint(lo, hi)
1012 }
1013
1014 pub fn quantile(&self, q: f64) -> f64 {
1015 let ends = |q: f64| {
1016 self.parts.iter().map(move |(part, _)| match part {
1017 Part::Point(x) => *x,
1018 Part::Continuous(f) => f.quantile(q),
1019 Part::Analytic(a) => a.quantile(q),
1020 })
1021 };
1022 if self.parts.len() == 1 {
1023 return match &self.parts[0].0 {
1024 Part::Point(x) => *x,
1025 Part::Continuous(f) => f.quantile(q),
1026 Part::Analytic(a) => a.quantile(q),
1027 };
1028 }
1029 if q <= 0.0 {
1030 return ends(0.0).fold(f64::INFINITY, f64::min);
1031 }
1032 if q >= 1.0 {
1033 return ends(1.0).fold(f64::NEG_INFINITY, f64::max);
1034 }
1035 let lo = ends(q.min(1e-12)).fold(f64::INFINITY, f64::min).max(-f64::MAX);
1036 let hi = ends(q.max(1.0 - 1e-12)).fold(f64::NEG_INFINITY, f64::max).min(f64::MAX);
1037 if lo >= hi {
1038 return lo;
1039 }
1040 if self.cdf(hi) < q {
1041 return f64::INFINITY;
1042 }
1043 invert(|x| self.cdf(x), q, (lo - 1e-9 * lo.abs().max(1.0)).max(-f64::MAX), hi)
1044 }
1045}
1046
1047#[cfg(test)]
1048mod tests {
1049 use super::*;
1050
1051 fn close(a: f64, b: f64, tol: f64) {
1052 assert!((a - b).abs() <= tol * (1.0 + b.abs()), "{a} vs {b}");
1053 }
1054
1055 #[test]
1056 fn normal_quantiles() {
1057 close(std_normal_cdf(Z95), 0.95, 1e-15);
1058 close(std_normal_quantile(0.95), Z95, 1e-14);
1059 close(std_normal_quantile(0.5), 0.0, 1e-15);
1060 close(std_normal_quantile(0.025), -1.9599639845400545, 1e-13);
1061 close(std_normal_quantile(1e-10), -6.361340902404056, 1e-10);
1062 close(std_normal_cdf(-1.959963984540054), 0.025, 1e-14);
1063 }
1064
1065 #[test]
1066 fn estimates_have_the_right_intervals() {
1067 let e = Family::estimate(3.0, 7.0).unwrap();
1068 close(e.quantile(0.05), 3.0, 1e-12);
1069 close(e.quantile(0.95), 7.0, 1e-12);
1070 let r = Family::normal_range(-0.08, 0.0).unwrap();
1071 close(r.quantile(0.05), -0.08, 1e-12);
1072 close(r.cdf(0.0), 0.95, 1e-12);
1073 assert!(Family::estimate(-1.0, 3.0).is_err());
1074 assert!(Family::estimate(3.0, 1.0).is_err());
1075 }
1076
1077 #[test]
1078 fn incomplete_functions() {
1079 close(beta_cdf(2.0, 3.0, 0.5), 11.0 / 16.0, 1e-14);
1081 close(beta_cdf(1.0, 1.0, 0.3), 0.3, 1e-14);
1082 close(gamma_cdf(1.0, 2.0), 1.0 - (-2.0f64).exp(), 1e-14);
1083 close(gamma_cdf(0.5, 3.0), libm::erf(3.0f64.sqrt()), 1e-13);
1084 close(gamma_cdf(10.0, 30.0), 0.9999928782491372, 1e-13);
1085 let b = Family::beta(2.0, 40.0).unwrap();
1086 close(b.quantile(b.cdf(0.05)), 0.05, 1e-10);
1087 let g = Family::gamma(3.0, 2.0).unwrap();
1088 close(g.quantile(0.5), 5.348120627447122, 1e-10);
1089 }
1090
1091 #[test]
1092 fn densities_integrate_to_their_cdfs() {
1093 let families = [
1094 Family::normal(1.0, 2.0).unwrap(),
1095 Family::lognormal(0.5, 0.3).unwrap(),
1096 Family::uniform(-1.0, 3.0).unwrap(),
1097 Family::beta(2.5, 4.0).unwrap(),
1098 Family::gamma(2.0, 1.5).unwrap(),
1099 Family::exponential(0.7).unwrap(),
1100 Family::triangular(0.0, 1.0, 4.0).unwrap(),
1101 Family::pert(1.0, 2.0, 6.0).unwrap(),
1102 ];
1103 for f in families {
1104 let (a, b) = (f.quantile(0.1), f.quantile(0.8));
1105 let n = 20_000;
1106 let h = (b - a) / n as f64;
1107 let integral: f64 = (0..n).map(|i| f.pdf(a + (i as f64 + 0.5) * h) * h).sum();
1108 close(integral, 0.7, 1e-6);
1109 close(f.cdf(b) - f.cdf(a), 0.7, 1e-9);
1110 }
1111 }
1112
1113 #[test]
1115 fn draws_follow_the_distributions() {
1116 let families = [
1117 Family::normal(1.0, 2.0).unwrap(),
1118 Family::lognormal(0.5, 0.3).unwrap(),
1119 Family::uniform(-1.0, 3.0).unwrap(),
1120 Family::beta(2.0, 40.0).unwrap(),
1121 Family::beta(0.5, 0.5).unwrap(),
1122 Family::gamma(0.4, 1.5).unwrap(),
1123 Family::gamma(7.0, 0.5).unwrap(),
1124 Family::exponential(0.7).unwrap(),
1125 Family::triangular(0.0, 1.0, 4.0).unwrap(),
1126 Family::pert(1.0, 2.0, 6.0).unwrap(),
1127 Family::estimate(60.0, 150.0).unwrap(),
1128 ];
1129 let mut rng = Rng::new(42);
1130 let n = 20_000;
1131 for f in families {
1132 let mut xs: Vec<f64> = (0..n).map(|_| f.sample(&mut rng)).collect();
1133 xs.sort_by(f64::total_cmp);
1134 let d = xs
1135 .iter()
1136 .enumerate()
1137 .map(|(i, x)| {
1138 let c = f.cdf(*x);
1139 (c - i as f64 / n as f64)
1140 .abs()
1141 .max((c - (i + 1) as f64 / n as f64).abs())
1142 })
1143 .fold(0.0, f64::max);
1144 assert!(d < 1.95 / (n as f64).sqrt(), "{f}: D = {d}");
1146 let mean = xs.iter().sum::<f64>() / n as f64;
1147 close(mean, f.mean(), 6.0 * f.variance().sqrt() / (n as f64).sqrt());
1148 }
1149 }
1150
1151 #[test]
1152 fn batches_have_streams_of_their_own() {
1153 let firsts: Vec<u64> = (0..1000).map(|i| Rng::stream(11, i).next_u64()).collect();
1154 let mut distinct = firsts.clone();
1155 distinct.sort();
1156 distinct.dedup();
1157 assert_eq!(distinct.len(), firsts.len());
1158 assert_eq!(Rng::stream(11, 3).next_u64(), Rng::stream(11, 3).next_u64());
1159 assert_ne!(Rng::stream(11, 3).next_u64(), Rng::stream(12, 3).next_u64());
1160 }
1161
1162 #[test]
1163 fn seeds_repeat() {
1164 let (mut a, mut b) = (Rng::new(7), Rng::new(7));
1165 for _ in 0..100 {
1166 assert_eq!(a.next_u64(), b.next_u64());
1167 }
1168 assert_ne!(Rng::new(7).next_u64(), Rng::new(8).next_u64());
1169 }
1170
1171 #[test]
1172 fn mixtures() {
1173 let m = Mixture {
1174 parts: vec![
1175 (Part::Point(0.0), 0.5),
1176 (Part::Continuous(Family::uniform(1.0, 3.0).unwrap()), 0.5),
1177 ],
1178 };
1179 close(m.mean(), 1.0, 1e-12);
1180 close(m.cdf(2.0), 0.75, 1e-12);
1181 close(m.quantile(0.75), 2.0, 1e-9);
1182 close(m.quantile(0.25), 0.0, 1e-9);
1183 }
1184}