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::{Fault, 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 {
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
68/// a/(a+b) for positive parameters, even when their sum overflows.
69fn 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    /// `a to b`: the lognormal whose 5% and 95% quantiles are `a` and `b`.
170    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    /// `normal_range(lo, hi)`: the normal whose 5% and 95% quantiles are
184    /// `lo` and `hi`.
185    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    /// The parameters of the beta distribution behind a PERT.
197    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    /// The smallest and largest possible values.
203    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    /// Conditional moments on a nonempty interval, using incomplete moments.
269    /// Uniform and normal use centered formulas to avoid subtracting large
270    /// location parameters when computing the variance.
271    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    /// P(X ≤ x).
408    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    /// The value below which a share `p` of the distribution lies.
450    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    /// One draw.
483    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    /// The parameters, for hashing, ordering and display.
513    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
542// ── Special functions ────────────────────────────────────────────────────
543
544fn 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
552/// Φ⁻¹(p): a rational approximation (Abramowitz and Stegun 26.2.23, error
553/// below 4.5e-4), polished by Newton's method on the exact CDF.
554pub 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        // At an end, the density is infinite, 1/B(a, b), or zero, as the
590        // parameter for that end is below, at or above 1.
591        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
607/// ln √(2π).
608const LN_SQRT_2PI: f64 = 0.918_938_533_204_672_7;
609
610/// ln B(a, b), the logarithm of the beta function. For large arguments,
611/// the log gammas it's made of are huge and nearly cancel, so it's computed
612/// from their Stirling series instead, as R's `lbeta` is.
613pub 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
627/// ln Γ(x) minus Stirling's approximation (x − ½) ln x − x + ln √(2π), for
628/// x ≥ 10: the first seven terms of its series, which are exact to 10⁻¹⁶
629/// there.
630fn 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
638/// The regularized incomplete beta function I_x(a, b), by its continued
639/// fraction (Numerical Recipes, section 6.4).
640pub 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
700/// The regularized lower incomplete gamma function P(a, x): its series
701/// below a + 1, its continued fraction above (Numerical Recipes, 6.2).
702pub 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
750/// The x in `lo..hi` where a non-decreasing `f` reaches `p` (the smallest
751/// such x), by bisection.
752pub 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// ── Random numbers ───────────────────────────────────────────────────────
775
776/// xoshiro256++, seeded through SplitMix64.
777#[derive(Clone, Debug)]
778pub struct Rng {
779    s: [u64; 4],
780}
781
782/// SplitMix64's step and output function.
783const 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    /// The random numbers of the batch numbered `index`, for a run with this
804    /// seed. Every batch has a stream of its own, so batches can run in any
805    /// order and on any number of threads, and still get the same numbers
806    /// (docs/semantics.md, section 14).
807    pub fn stream(seed: u64, index: u64) -> Rng {
808        // The index-th output of SplitMix64 from the seed: different for
809        // every index.
810        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    /// Uniform in [0, 1).
827    pub fn uniform(&mut self) -> f64 {
828        (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
829    }
830
831    /// Uniform in (0, 1]: safe to take the logarithm of.
832    pub fn open(&mut self) -> f64 {
833        1.0 - self.uniform()
834    }
835
836    /// A standard normal draw (Box–Muller).
837    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    /// A gamma draw with scale 1 (Marsaglia and Tsang, 2000).
843    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        // Both draws round to zero only for tiny parameters, where the
879        // distribution is nearly all at 0 and 1.
880        if self.uniform() < ratio(a, b) { 1.0 } else { 0.0 }
881    }
882
883    /// An index chosen with probability proportional to `weights`, or `None`
884    /// if they're all zero.
885    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// ── Mixtures ─────────────────────────────────────────────────────────────
908
909/// A part of a mixture: a number, or a continuous distribution.
910#[derive(Clone, Debug)]
911pub enum Part {
912    Point(f64),
913    Continuous(Family),
914    Analytic(crate::analytic::Analytic),
915}
916
917/// Moments, CDF and quantiles of a mixture of numbers and continuous
918/// distributions, each with its probability.
919#[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    /// Distinguish an actual gap in the support at half the mass from a CDF
976    /// that merely rounds to 0.5. This also handles disconnected analytic
977    /// domains in report summaries without perturbing the requested quantile.
978    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        // I_0.5(2, 3) = 11/16; P(1, x) = 1 − e^−x; P(1/2, x) = erf(√x).
1086        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    /// Kolmogorov–Smirnov: the draws follow each distribution's CDF.
1120    #[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            // The 0.001 critical value is 1.95 / √n.
1151            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}