Skip to main content

probl_engine/
continuous.rs

1//! Continuous distributions (docs/semantics.md, section 13): densities,
2//! CDFs, quantiles and moments, and drawing from them.
3//!
4//! Drawing uses its own generator and `libm` rather than the platform's math
5//! library, so that a seed gives the same numbers everywhere (section 14).
6
7use crate::error::{OpError, OpResult};
8use std::f64::consts::{PI, SQRT_2};
9use std::fmt;
10
11/// The 95% quantile of the standard normal: an estimate's 90% interval is
12/// this many standard deviations on each side of its centre.
13pub const Z95: f64 = 1.644853626951472;
14
15/// A continuous distribution.
16#[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    /// A beta distribution stretched over `lo..hi` with its mode at `mode`.
47    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
64/// a/(a+b) for positive parameters, even when their sum overflows.
65fn 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    /// `a to b`: the lognormal whose 5% and 95% quantiles are `a` and `b`.
166    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    /// `normal_range(lo, hi)`: the normal whose 5% and 95% quantiles are
178    /// `lo` and `hi`.
179    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    /// The parameters of the beta distribution behind a PERT.
191    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    /// The smallest and largest possible values.
197    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    /// Conditional moments on a nonempty interval, using incomplete moments.
263    /// Uniform and normal use centered formulas to avoid subtracting large
264    /// location parameters when computing the variance.
265    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    /// P(X ≤ x).
402    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    /// The value below which a share `p` of the distribution lies.
444    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    /// One draw.
477    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    /// The parameters, for hashing, ordering and display.
507    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
536// ── Special functions ────────────────────────────────────────────────────
537
538fn 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
546/// Φ⁻¹(p): a rational approximation (Abramowitz and Stegun 26.2.23, error
547/// below 4.5e-4), polished by Newton's method on the exact CDF.
548pub 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        // At an end, the density is infinite, 1/B(a, b), or zero, as the
584        // parameter for that end is below, at or above 1.
585        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
601/// ln √(2π).
602const LN_SQRT_2PI: f64 = 0.918_938_533_204_672_7;
603
604/// ln B(a, b), the logarithm of the beta function. For large arguments,
605/// the log gammas it's made of are huge and nearly cancel, so it's computed
606/// from their Stirling series instead, as R's `lbeta` is.
607pub 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
621/// ln Γ(x) minus Stirling's approximation (x − ½) ln x − x + ln √(2π), for
622/// x ≥ 10: the first seven terms of its series, which are exact to 10⁻¹⁶
623/// there.
624fn 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
632/// The regularized incomplete beta function I_x(a, b), by its continued
633/// fraction (Numerical Recipes, section 6.4).
634pub 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
694/// The regularized lower incomplete gamma function P(a, x): its series
695/// below a + 1, its continued fraction above (Numerical Recipes, 6.2).
696pub 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
744/// The x in `lo..hi` where a non-decreasing `f` reaches `p` (the smallest
745/// such x), by bisection.
746pub 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// ── Random numbers ───────────────────────────────────────────────────────
769
770/// xoshiro256++, seeded through SplitMix64.
771#[derive(Clone, Debug)]
772pub struct Rng {
773    s: [u64; 4],
774}
775
776/// SplitMix64's step and output function.
777const 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    /// The random numbers of the batch numbered `index`, for a run with this
798    /// seed. Every batch has a stream of its own, so batches can run in any
799    /// order and on any number of threads, and still get the same numbers
800    /// (docs/semantics.md, section 14).
801    pub fn stream(seed: u64, index: u64) -> Rng {
802        // The index-th output of SplitMix64 from the seed: different for
803        // every index.
804        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    /// Uniform in [0, 1).
821    pub fn uniform(&mut self) -> f64 {
822        (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
823    }
824
825    /// Uniform in (0, 1]: safe to take the logarithm of.
826    pub fn open(&mut self) -> f64 {
827        1.0 - self.uniform()
828    }
829
830    /// A standard normal draw (Box–Muller).
831    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    /// A gamma draw with scale 1 (Marsaglia and Tsang, 2000).
837    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        // Both draws round to zero only for tiny parameters, where the
873        // distribution is nearly all at 0 and 1.
874        if self.uniform() < ratio(a, b) { 1.0 } else { 0.0 }
875    }
876
877    /// An index chosen with probability proportional to `weights`, or `None`
878    /// if they're all zero.
879    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// ── Mixtures ─────────────────────────────────────────────────────────────
902
903/// A part of a mixture: a number, or a continuous distribution.
904#[derive(Clone, Debug)]
905pub enum Part {
906    Point(f64),
907    Continuous(Family),
908    Analytic(crate::analytic::Analytic),
909}
910
911/// Moments, CDF and quantiles of a mixture of numbers and continuous
912/// distributions, each with its probability.
913#[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    /// Distinguish an actual gap in the support at half the mass from a CDF
970    /// that merely rounds to 0.5. This also handles disconnected analytic
971    /// domains in report summaries without perturbing the requested quantile.
972    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        // I_0.5(2, 3) = 11/16; P(1, x) = 1 − e^−x; P(1/2, x) = erf(√x).
1080        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    /// Kolmogorov–Smirnov: the draws follow each distribution's CDF.
1114    #[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            // The 0.001 critical value is 1.95 / √n.
1145            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}