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