Skip to main content

probl_engine/
dist.rs

1//! Finite distributions: dice, choices, counts, and the results of computing
2//! with them.
3
4use crate::continuous::Rng;
5use crate::error::{OpError, OpResult};
6use crate::value::Value;
7use std::cmp::Ordering;
8use std::collections::BTreeMap;
9use std::hash::{Hash, Hasher};
10use std::sync::Arc;
11use std::sync::atomic::{AtomicBool, AtomicU64, Ordering as Atomic};
12
13/// Outcomes below this probability are dropped from infinite supports.
14const TAIL: f64 = 1e-18;
15
16#[derive(Clone, Debug)]
17pub struct Dist {
18    /// Sorted by value, each value once, every weight positive.
19    pub outcomes: Vec<(Value, f64)>,
20    /// Probability not represented by `outcomes`: a tail cut off to keep the
21    /// support finite, or weight a `simulate` left unresolved. Drawing counts
22    /// it as unresolved.
23    pub missing: f64,
24}
25
26/// Work taken from a shared budget at a time: threads rarely touch it, and
27/// between them hold back little of it.
28const WORK_CHUNK: u64 = 1 << 14;
29
30/// Limits on building distributions, shared by everything in a run.
31#[derive(Clone, Debug)]
32pub struct Budget {
33    /// Maximum bits in the magnitude of a single integer.
34    pub max_integer_bits: u64,
35    /// Cumulative allowance for produced large integer payloads, shared by runs.
36    pub integer_bytes_left: Arc<AtomicU64>,
37    pub max_string_bytes: usize,
38    /// Cumulative string payload allowance, shared by sampled workers.
39    pub string_bytes_left: Arc<AtomicU64>,
40    pub cancel: Option<Arc<AtomicBool>>,
41    /// The most outcomes one distribution (or combination) may have.
42    pub max_outcomes: usize,
43    /// The most elements a collection built by the program may have.
44    pub max_collection: usize,
45    /// Units of work left: world-steps plus outcomes computed.
46    pub work_left: u64,
47    /// Work shared with other threads, which `work_left` is topped up from
48    /// (when batches of runs are sampled in parallel).
49    pub shared: Option<Arc<AtomicU64>>,
50}
51
52impl Budget {
53    pub fn unlimited() -> Budget {
54        Budget {
55            max_integer_bits: probl_number::MAX_INTEGER_BITS,
56            integer_bytes_left: Arc::new(AtomicU64::new(u64::MAX)),
57            max_string_bytes: usize::MAX,
58            string_bytes_left: Arc::new(AtomicU64::new(u64::MAX)),
59            cancel: None,
60            max_outcomes: usize::MAX,
61            max_collection: usize::MAX,
62            work_left: u64::MAX,
63            shared: None,
64        }
65    }
66
67    pub fn integer_bits(&self, bits: u64) -> OpResult<()> {
68        let limit = self.max_integer_bits.min(probl_number::MAX_INTEGER_BITS);
69        if bits > limit {
70            return Err(OpError::limit(format!(
71                "integer size exceeds the limit of {limit} bits"
72            )));
73        }
74        Ok(())
75    }
76
77    pub fn string_size(&self, bytes: usize) -> OpResult<()> {
78        if bytes > self.max_string_bytes {
79            return Err(OpError::limit(format!(
80                "string size exceeds the limit of {} UTF-8 bytes",
81                self.max_string_bytes
82            )));
83        }
84        Ok(())
85    }
86
87    /// Charge a scan before doing it, in units of 64 UTF-8 bytes.
88    pub fn string_work(&mut self, s: &str) -> OpResult<()> {
89        self.string_size(s.len())?;
90        self.work((s.len() as u64).div_ceil(64).max(1))
91    }
92
93    /// Reserve payload bytes before allocation or growing a text builder.
94    // `try_update`, its new name, needs a newer Rust than the crates support (1.85).
95    #[allow(deprecated)]
96    pub fn string_allocation(&self, bytes: usize) -> OpResult<()> {
97        self.string_size(bytes)?;
98        self.string_bytes_left
99            .fetch_update(Atomic::Relaxed, Atomic::Relaxed, |left| left.checked_sub(bytes as u64))
100            .map(|_| ())
101            .map_err(|_| OpError::limit("the run used up its string memory allowance"))
102    }
103
104    /// Reserve before materializing a collection, or after a single bounded
105    /// result. Conservative: shared results may be charged again; no refunds.
106    // `try_update`, its new name, needs a newer Rust than the crates support (1.85).
107    #[allow(deprecated)]
108    pub fn integer_allocation(&self, bits: u64, count: u64) -> OpResult<()> {
109        self.integer_bits(bits)?;
110        if bits <= 63 || count == 0 {
111            return Ok(());
112        }
113        let bytes = bits
114            .div_ceil(64)
115            .saturating_mul(8)
116            .saturating_add(48)
117            .saturating_mul(count);
118        self.integer_bytes_left
119            .fetch_update(Atomic::Relaxed, Atomic::Relaxed, |left| left.checked_sub(bytes))
120            .map(|_| ())
121            .map_err(|_| OpError::limit("the run used up its large integer memory allowance"))
122    }
123
124    /// Charge before arithmetic, including copying and processing large operands.
125    pub fn integer_work(
126        &mut self,
127        a: &probl_number::Integer,
128        b: &probl_number::Integer,
129        quadratic: bool,
130    ) -> OpResult<()> {
131        self.integer_bits(a.bits())?;
132        self.integer_bits(b.bits())?;
133        let x = a.bits().div_ceil(64).max(1);
134        let y = b.bits().div_ceil(64).max(1);
135        self.work(if quadratic { x.saturating_mul(y) } else { x.max(y) })
136    }
137
138    /// Check that a distribution with `n` outcomes may be built.
139    pub fn outcomes(&self, n: u128) -> OpResult<()> {
140        if n > self.max_outcomes as u128 {
141            let shown = if n > 1_000_000_000_000 {
142                "more than 10¹²".to_string()
143            } else {
144                n.to_string()
145            };
146            return Err(OpError::limit(format!(
147                "a distribution with {shown} outcomes is over the limit of {}",
148                self.max_outcomes
149            )));
150        }
151        Ok(())
152    }
153
154    /// Check that a collection with `n` elements may be built.
155    pub fn collection(&self, n: u128) -> OpResult<()> {
156        if n > self.max_collection as u128 {
157            return Err(OpError::limit(format!(
158                "a collection with {n} elements is over the limit of {}",
159                self.max_collection
160            )));
161        }
162        Ok(())
163    }
164
165    /// Spend `n` units of work.
166    pub fn work(&mut self, n: u64) -> OpResult<()> {
167        if self.cancel.as_ref().is_some_and(|c| c.load(Atomic::Relaxed)) {
168            return Err(OpError::limit("the run was cancelled"));
169        }
170        if n > self.work_left && !self.top_up(n - self.work_left) {
171            self.work_left = 0;
172            return Err(OpError::limit("the run used up its work budget"));
173        }
174        self.work_left -= n;
175        Ok(())
176    }
177
178    /// Take at least `need` units from the shared budget, if it has them.
179    // `try_update`, its new name, needs a newer Rust than the crates support (1.85).
180    #[allow(deprecated)]
181    fn top_up(&mut self, need: u64) -> bool {
182        let Some(shared) = &self.shared else {
183            return false;
184        };
185        let want = need.max(WORK_CHUNK);
186        let taken = shared.fetch_update(Atomic::Relaxed, Atomic::Relaxed, |left| {
187            (left >= need).then(|| left - want.min(left))
188        });
189        match taken {
190            Ok(left) => {
191                self.work_left += want.min(left);
192                true
193            }
194            Err(_) => false,
195        }
196    }
197
198    /// Give the work not spent back to the shared budget.
199    pub fn give_back(&mut self) {
200        if let Some(shared) = &self.shared {
201            shared.fetch_add(std::mem::take(&mut self.work_left), Atomic::Relaxed);
202        }
203    }
204}
205
206impl Dist {
207    pub fn point(v: Value) -> Dist {
208        Dist {
209            outcomes: vec![(v, 1.0)],
210            missing: 0.0,
211        }
212    }
213
214    /// Build from (value, weight) pairs in any order, whose weights add up
215    /// to `1 - missing`; equal values are merged.
216    ///
217    /// The weights are rescaled to add up to exactly that, which removes the
218    /// rounding of the sums that produced them: six sixths of `d6 > 0` are
219    /// 0.9999999999999999, but the distribution is certainly `true`. Without
220    /// this, a certain condition would leave a branch of weight 1e-16.
221    pub fn from_pairs(mut pairs: Vec<(Value, f64)>, missing: f64) -> Dist {
222        pairs.retain(|(_, w)| *w > 0.0);
223        pairs.sort_by(|a, b| a.0.cmp(&b.0));
224        let mut outcomes: Vec<(Value, f64)> = Vec::with_capacity(pairs.len());
225        for (v, w) in pairs {
226            match outcomes.last_mut() {
227                Some((last, total)) if *last == v => *total += w,
228                _ => outcomes.push((v, w)),
229            }
230        }
231        Dist::normalized(outcomes, missing)
232    }
233
234    /// Like `from_pairs`, for pairs already sorted by value, each value once.
235    fn from_sorted(mut pairs: Vec<(Value, f64)>, missing: f64) -> Dist {
236        debug_assert!(pairs.windows(2).all(|p| p[0].0 < p[1].0), "unsorted outcomes");
237        pairs.retain(|(_, w)| *w > 0.0);
238        Dist::normalized(pairs, missing)
239    }
240
241    fn normalized(mut outcomes: Vec<(Value, f64)>, missing: f64) -> Dist {
242        let target = 1.0 - missing;
243        let total = crate::stats::sum(outcomes.iter().map(|(_, w)| *w));
244        if let [(_, w)] = outcomes.as_mut_slice() {
245            *w = target;
246        } else if target > 0.0 && total > 0.0 && total != target {
247            let scale = target / total;
248            for (_, w) in &mut outcomes {
249                *w *= scale;
250            }
251        }
252        Dist { outcomes, missing }
253    }
254
255    pub fn into_value(self) -> Value {
256        Value::Dist(Arc::new(self))
257    }
258
259    pub fn uniform(values: Vec<Value>) -> Dist {
260        let p = 1.0 / values.len() as f64;
261        Dist::from_pairs(values.into_iter().map(|v| (v, p)).collect(), 0.0)
262    }
263
264    /// `true` with probability `p`.
265    pub fn bernoulli(p: f64) -> Dist {
266        Dist::from_sorted(vec![(Value::Bool(false), 1.0 - p), (Value::Bool(true), p)], 0.0)
267    }
268
269    /// The sum of `count` dice with `sides` sides.
270    pub fn dice(count: u32, sides: u32, budget: &mut Budget) -> OpResult<Dist> {
271        let support = count as u128 * (sides as u128 - 1) + 1;
272        budget.outcomes(support)?;
273        budget.work((count as u128 * support * sides as u128).min(u64::MAX as u128) as u64)?;
274        let mut sums = vec![1.0];
275        let p = 1.0 / sides as f64;
276        for _ in 0..count {
277            let mut next = vec![0.0; sums.len() + sides as usize];
278            for (s, w) in sums.iter().enumerate() {
279                if *w == 0.0 {
280                    continue;
281                }
282                for face in 1..=sides as usize {
283                    next[s + face] += w * p;
284                }
285            }
286            sums = next;
287        }
288        let pairs = sums
289            .into_iter()
290            .enumerate()
291            .filter(|(_, w)| *w > 0.0)
292            .map(|(s, w)| (Value::Int((s as i64).into()), w))
293            .collect();
294        Ok(Dist::from_sorted(pairs, 0.0))
295    }
296
297    pub fn binomial(n: u64, p: f64, budget: &mut Budget) -> OpResult<Dist> {
298        if p <= 0.0 {
299            return Ok(Dist::point(Value::Int(0.into())));
300        }
301        if p >= 1.0 {
302            return Ok(Dist::point(Value::Int((n as i64).into())));
303        }
304        let odds = p / (1.0 - p);
305        let mode = (((n as f64) + 1.0) * p).floor().min(n as f64) as u64;
306        walk_from_mode(
307            mode,
308            n,
309            |k| (n - k) as f64 / (k + 1) as f64 * odds,
310            |k| k as f64 / (n - k + 1) as f64 / odds,
311            budget,
312        )
313    }
314
315    pub fn poisson(rate: f64, budget: &mut Budget) -> OpResult<Dist> {
316        if rate <= 0.0 {
317            return Ok(Dist::point(Value::Int(0.into())));
318        }
319        walk_from_mode(
320            rate.floor() as u64,
321            u64::MAX,
322            |k| rate / (k + 1) as f64,
323            |k| k as f64 / rate,
324            budget,
325        )
326    }
327
328    /// Number of tries up to and including the first success.
329    pub fn geometric(p: f64, budget: &mut Budget) -> OpResult<Dist> {
330        if p >= 1.0 {
331            return Ok(Dist::point(Value::Int(1.into())));
332        }
333        // Outcomes until the tail (1 - p)^k drops below the threshold.
334        let count = (libm::log(TAIL) / libm::log(1.0 - p)).ceil();
335        budget.outcomes(if count.is_finite() { count as u128 } else { u128::MAX })?;
336        budget.work(count as u64)?;
337        let mut pairs = Vec::with_capacity(count as usize);
338        let mut tail = 1.0;
339        let mut k: i64 = 1;
340        while tail >= TAIL {
341            pairs.push((Value::Int(k.into()), tail * p));
342            tail *= 1.0 - p;
343            k += 1;
344        }
345        Ok(Dist::from_sorted(pairs, tail))
346    }
347
348    /// The dice of `count` rolls of `die`, sorted from highest to lowest.
349    pub fn pool(count: u32, die: &Dist, budget: &mut Budget) -> OpResult<Dist> {
350        // As many pools as multisets of `count` faces.
351        let faces = die.outcomes.len() as u128;
352        budget.outcomes(multisets(faces, count as u128))?;
353        let mut pools: BTreeMap<Vec<Value>, f64> = BTreeMap::new();
354        pools.insert(Vec::new(), 1.0);
355        for _ in 0..count {
356            budget.work(pools.len() as u64 * faces as u64)?;
357            let mut next: BTreeMap<Vec<Value>, f64> = BTreeMap::new();
358            for (pool, w) in &pools {
359                for (face, p) in &die.outcomes {
360                    let mut grown = pool.clone();
361                    let at = grown.iter().position(|v| v < face).unwrap_or(grown.len());
362                    grown.insert(at, face.clone());
363                    *next.entry(grown).or_insert(0.0) += w * p;
364                }
365            }
366            pools = next;
367        }
368        let pairs = pools.into_iter().map(|(pool, w)| (Value::list(pool), w)).collect();
369        // Each die independently misses `die.missing` of its probability.
370        let missing = -libm::expm1(count as f64 * libm::log1p(-die.missing));
371        Ok(Dist::from_pairs(pairs, missing))
372    }
373
374    pub fn total(&self) -> f64 {
375        crate::stats::sum(self.outcomes.iter().map(|(_, w)| *w))
376    }
377
378    /// For a distribution of facts: the probabilities of `true` and `false`.
379    pub fn truth(&self) -> Option<(f64, f64)> {
380        let (mut yes, mut no) = (0.0, 0.0);
381        for (v, w) in &self.outcomes {
382            match v {
383                Value::Bool(true) => yes += w,
384                Value::Bool(false) => no += w,
385                _ => return None,
386            }
387        }
388        Some((yes, no))
389    }
390
391    fn numbers(&self) -> Option<Vec<(f64, f64)>> {
392        self.outcomes
393            .iter()
394            .map(|(v, w)| match v {
395                Value::Bool(_) => None,
396                _ => v.as_f64().map(|x| (x, *w)),
397            })
398            .collect()
399    }
400
401    /// The mean of the resolved outcomes.
402    pub fn mean(&self) -> Option<f64> {
403        let nums = self.numbers()?;
404        Some(crate::stats::weighted_mean(nums.iter().copied()))
405    }
406
407    pub fn variance(&self) -> Option<f64> {
408        let sd = self.sd()?;
409        Some(sd * sd)
410    }
411
412    pub fn sd(&self) -> Option<f64> {
413        let nums = self.numbers()?;
414        Some(crate::stats::weighted_sd(nums.iter().map(|(x, w)| (*x, 0.0, *w))))
415    }
416
417    /// Quantile in storage order. Language queries first build a population in
418    /// language order, which can differ for nested mixed numeric values.
419    pub fn quantile(&self, q: f64) -> Option<Value> {
420        crate::stats::quantile(&self.outcomes, q).cloned()
421    }
422}
423
424/// `binomial`, `poisson` or `geometric` by their parameters: their
425/// probabilities and draws, without listing their outcomes. Sampling uses
426/// them (docs/semantics.md, section 14).
427#[derive(Clone, Copy, Debug)]
428pub enum Counts {
429    Binomial {
430        n: u64,
431        p: f64,
432    },
433    Poisson {
434        rate: f64,
435    },
436    /// Tries up to and including the first success.
437    Geometric {
438        p: f64,
439    },
440}
441
442impl Counts {
443    /// Whether a draw can walk out from the mode: the spread is small enough
444    /// for that to take a bounded number of steps. Otherwise listing the
445    /// outcomes, which is budgeted, fails cleanly.
446    pub fn direct(&self) -> bool {
447        match *self {
448            Counts::Binomial { n, p } => (n as f64) * p * (1.0 - p) <= 1e10,
449            Counts::Poisson { rate } => rate <= 1e10,
450            Counts::Geometric { .. } => true,
451        }
452    }
453
454    /// The distribution with its outcomes listed.
455    pub fn list(&self, budget: &mut Budget) -> OpResult<Dist> {
456        match *self {
457            Counts::Binomial { n, p } => Dist::binomial(n, p, budget),
458            Counts::Poisson { rate } => Dist::poisson(rate, budget),
459            Counts::Geometric { p } => Dist::geometric(p, budget),
460        }
461    }
462
463    /// P(X = k).
464    pub fn pmf(&self, k: f64) -> f64 {
465        if k < 0.0 || k.fract() != 0.0 {
466            return 0.0;
467        }
468        let exactly = |x: f64| if k == x { 1.0 } else { 0.0 };
469        match *self {
470            Counts::Binomial { n, p } => {
471                let n = n as f64;
472                if k > n {
473                    0.0
474                } else if p <= 0.0 {
475                    exactly(0.0)
476                } else if p >= 1.0 {
477                    exactly(n)
478                } else {
479                    crate::math::exp(
480                        libm::lgamma(n + 1.0) - libm::lgamma(k + 1.0) - libm::lgamma(n - k + 1.0)
481                            + k * libm::log(p)
482                            + (n - k) * libm::log1p(-p),
483                    )
484                }
485            }
486            Counts::Poisson { rate } => {
487                if rate <= 0.0 {
488                    exactly(0.0)
489                } else {
490                    crate::math::exp(k * libm::log(rate) - rate - libm::lgamma(k + 1.0))
491                }
492            }
493            Counts::Geometric { p } => {
494                if k < 1.0 {
495                    0.0
496                } else if p >= 1.0 {
497                    exactly(1.0)
498                } else {
499                    p * crate::math::exp((k - 1.0) * libm::log1p(-p))
500                }
501            }
502        }
503    }
504
505    /// One draw.
506    pub fn sample(&self, rng: &mut Rng) -> i64 {
507        match *self {
508            Counts::Binomial { n, p } => {
509                if p <= 0.0 {
510                    return 0;
511                }
512                if p >= 1.0 {
513                    return n as i64;
514                }
515                let odds = p / (1.0 - p);
516                let mode = (((n as f64) + 1.0) * p).floor().min(n as f64) as u64;
517                let up = |k: u64| (n - k) as f64 / (k + 1) as f64 * odds;
518                let down = |k: u64| k as f64 / (n - k + 1) as f64 / odds;
519                from_mode(rng, mode, self.pmf(mode as f64), n, up, down) as i64
520            }
521            Counts::Poisson { rate } => {
522                if rate <= 0.0 {
523                    return 0;
524                }
525                let mode = rate.floor() as u64;
526                let up = |k: u64| rate / (k + 1) as f64;
527                let down = |k: u64| k as f64 / rate;
528                from_mode(rng, mode, self.pmf(mode as f64), u64::MAX, up, down) as i64
529            }
530            Counts::Geometric { p } => {
531                if p >= 1.0 {
532                    return 1;
533                }
534                // The inverse of P(X ≤ k) = 1 − (1 − p)^k.
535                (libm::log(rng.open()) / libm::log1p(-p)).ceil().max(1.0) as i64
536            }
537        }
538    }
539}
540
541/// Inversion from the mode outward: a uniform draw is compared with the
542/// probabilities added up from the mode, taking the more likely neighbour
543/// next, so that a draw usually takes a few steps. `up` and `down` are as in
544/// `walk_from_mode`.
545fn from_mode(
546    rng: &mut Rng,
547    mode: u64,
548    p_mode: f64,
549    max: u64,
550    up: impl Fn(u64) -> f64,
551    down: impl Fn(u64) -> f64,
552) -> u64 {
553    let u = rng.uniform();
554    let mut total = p_mode;
555    let (mut lo, mut hi, mut p_lo, mut p_hi) = (mode, mode, p_mode, p_mode);
556    let mut last = mode;
557    while u >= total {
558        let below = if lo > 0 { p_lo * down(lo) } else { 0.0 };
559        let above = if hi < max { p_hi * up(hi) } else { 0.0 };
560        if below <= 0.0 && above <= 0.0 {
561            // The probabilities ran out before reaching the draw: that is
562            // only rounding.
563            break;
564        }
565        if above >= below {
566            hi += 1;
567            p_hi = above;
568            total += above;
569            last = hi;
570        } else {
571            lo -= 1;
572            p_lo = below;
573            total += below;
574            last = lo;
575        }
576    }
577    last
578}
579
580/// The outcomes `0..=max` of a count distribution with a single mode, built
581/// from the ratios between neighbouring probabilities: `up(k)` is
582/// P(k + 1) / P(k) and `down(k)` is P(k − 1) / P(k). Working with ratios
583/// relative to the mode stays accurate for huge parameters. The walk stops on
584/// each side once an outcome is below the threshold; the rest of that tail
585/// shrinks at least geometrically, which bounds the missing mass.
586fn walk_from_mode(
587    mode: u64,
588    max: u64,
589    up: impl Fn(u64) -> f64,
590    down: impl Fn(u64) -> f64,
591    budget: &mut Budget,
592) -> OpResult<Dist> {
593    let mut below: Vec<(u64, f64)> = Vec::new();
594    let mut above: Vec<(u64, f64)> = Vec::new();
595    let mut tails = 0.0;
596    let (mut k, mut w) = (mode, 1.0);
597    while k > 0 {
598        let next = w * down(k);
599        k -= 1;
600        if next < TAIL {
601            tails += geometric_tail(next, if k > 0 { down(k) } else { 0.0 });
602            break;
603        }
604        below.push((k, next));
605        w = next;
606        budget.outcomes((below.len() + above.len() + 1) as u128)?;
607        budget.work(1)?;
608    }
609    let (mut k, mut w) = (mode, 1.0);
610    while k < max {
611        let next = w * up(k);
612        k += 1;
613        if next < TAIL {
614            tails += geometric_tail(next, if k < max { up(k) } else { 0.0 });
615            break;
616        }
617        above.push((k, next));
618        w = next;
619        budget.outcomes((below.len() + above.len() + 1) as u128)?;
620        budget.work(1)?;
621    }
622    let sum = 1.0 + below.iter().map(|(_, w)| w).sum::<f64>() + above.iter().map(|(_, w)| w).sum::<f64>() + tails;
623    let pairs = below
624        .into_iter()
625        .rev()
626        .chain(std::iter::once((mode, 1.0)))
627        .chain(above)
628        .map(|(k, w)| (Value::Int((k as i64).into()), w / sum))
629        .collect();
630    Ok(Dist::from_sorted(pairs, tails / sum))
631}
632
633/// An upper bound on `first + first·r + first·r² + …` for ratios at most `r`.
634fn geometric_tail(first: f64, ratio: f64) -> f64 {
635    if first <= 0.0 {
636        return 0.0;
637    }
638    first / (1.0 - ratio.clamp(0.0, 1.0 - 1e-6))
639}
640
641/// The number of multisets of size `k` drawn from `n` kinds: C(n + k − 1, k),
642/// saturating.
643fn multisets(n: u128, k: u128) -> u128 {
644    if n == 0 {
645        return if k == 0 { 1 } else { 0 };
646    }
647    let mut result: u128 = 1;
648    for i in 1..=k {
649        result = result.saturating_mul(n + i - 1) / i;
650        if result > u64::MAX as u128 {
651            return u128::MAX;
652        }
653    }
654    result
655}
656
657impl PartialEq for Dist {
658    fn eq(&self, other: &Dist) -> bool {
659        self.missing.to_bits() == other.missing.to_bits()
660            && self.outcomes.len() == other.outcomes.len()
661            && self
662                .outcomes
663                .iter()
664                .zip(&other.outcomes)
665                .all(|((a, p), (b, q))| a == b && p.to_bits() == q.to_bits())
666    }
667}
668
669impl Eq for Dist {}
670
671impl Hash for Dist {
672    fn hash<H: Hasher>(&self, state: &mut H) {
673        self.outcomes.len().hash(state);
674        for (v, w) in &self.outcomes {
675            v.hash(state);
676            w.to_bits().hash(state);
677        }
678        self.missing.to_bits().hash(state);
679    }
680}
681
682impl Ord for Dist {
683    fn cmp(&self, other: &Dist) -> Ordering {
684        for ((a, p), (b, q)) in self.outcomes.iter().zip(&other.outcomes) {
685            let c = a.cmp(b).then_with(|| p.total_cmp(q));
686            if c != Ordering::Equal {
687                return c;
688            }
689        }
690        self.outcomes
691            .len()
692            .cmp(&other.outcomes.len())
693            .then_with(|| self.missing.total_cmp(&other.missing))
694    }
695}
696
697impl PartialOrd for Dist {
698    fn partial_cmp(&self, other: &Dist) -> Option<Ordering> {
699        Some(self.cmp(other))
700    }
701}
702
703/// Natural log of the gamma function (Lanczos approximation, g = 7).
704pub fn ln_gamma(x: f64) -> f64 {
705    const G: f64 = 7.0;
706    const C: [f64; 9] = [
707        0.999_999_999_999_809_9,
708        676.520_368_121_885_1,
709        -1_259.139_216_722_402_8,
710        771.323_428_777_653_1,
711        -176.615_029_162_140_6,
712        12.507_343_278_686_905,
713        -0.138_571_095_265_720_12,
714        9.984_369_578_019_572e-6,
715        1.505_632_735_149_311_6e-7,
716    ];
717    if x < 0.5 {
718        // Reflection formula.
719        return libm::log(std::f64::consts::PI / libm::sin(std::f64::consts::PI * x)) - ln_gamma(1.0 - x);
720    }
721    let x = x - 1.0;
722    let mut a = C[0];
723    let t = x + G + 0.5;
724    for (i, c) in C.iter().enumerate().skip(1) {
725        a += c / (x + i as f64);
726    }
727    0.5 * libm::log(2.0 * std::f64::consts::PI) + (x + 0.5) * libm::log(t) - t + libm::log(a)
728}
729
730#[cfg(test)]
731mod tests {
732    use super::*;
733
734    fn close(a: f64, b: f64) -> bool {
735        (a - b).abs() < 1e-12
736    }
737
738    #[test]
739    fn two_dice() {
740        let d = Dist::dice(2, 6, &mut Budget::unlimited()).unwrap();
741        assert_eq!(d.outcomes.len(), 11);
742        assert!(close(d.outcomes[5].1, 6.0 / 36.0));
743        assert!(close(d.mean().unwrap(), 7.0));
744        assert!(close(d.variance().unwrap(), 35.0 / 6.0));
745        assert_eq!(d.quantile(0.5), Some(Value::Int(7.into())));
746    }
747
748    #[test]
749    fn counts_sum_to_one() {
750        let b = &mut Budget::unlimited();
751        for d in [
752            Dist::binomial(30, 0.3, b).unwrap(),
753            Dist::poisson(4.5, b).unwrap(),
754            Dist::geometric(1.0 / 6.0, b).unwrap(),
755            Dist::binomial(30_000, 0.03, b).unwrap(),
756            Dist::poisson(1e9, b).unwrap(),
757        ] {
758            assert!((d.total() + d.missing - 1.0).abs() < 1e-9, "{:?}", d.outcomes.len());
759            assert!(d.missing < 1e-12);
760        }
761        assert!((Dist::poisson(4.5, b).unwrap().mean().unwrap() - 4.5).abs() < 1e-9);
762        assert!((Dist::binomial(30_000, 0.03, b).unwrap().mean().unwrap() - 900.0).abs() < 1e-6);
763    }
764
765    /// Drawing a count directly (as sampling does) follows the same
766    /// distribution as listing its outcomes: a chi-square test with the
767    /// cells of small expected counts pooled.
768    #[test]
769    fn direct_draws_follow_the_listed_distributions() {
770        let b = &mut Budget::unlimited();
771        let cases = [
772            Counts::Binomial { n: 250, p: 0.034 },
773            Counts::Binomial { n: 10, p: 0.5 },
774            Counts::Binomial { n: 5, p: 0.97 },
775            Counts::Binomial { n: 12_000, p: 0.03 },
776            Counts::Poisson { rate: 0.3 },
777            Counts::Poisson { rate: 100.5 },
778            Counts::Geometric { p: 0.2 },
779        ];
780        let mut rng = Rng::new(3);
781        for c in cases {
782            let d = c.list(b).unwrap();
783            for (v, p) in &d.outcomes {
784                let q = c.pmf(v.as_f64().unwrap());
785                assert!(
786                    (q - p).abs() <= 1e-10 * p.max(1e-300) + 1e-15,
787                    "{c:?} at {v}: {q} vs {p}"
788                );
789            }
790            let n = 200_000;
791            let mut seen: BTreeMap<i64, f64> = BTreeMap::new();
792            for _ in 0..n {
793                *seen.entry(c.sample(&mut rng)).or_default() += 1.0;
794            }
795            let (mut chi2, mut cells) = (0.0, 0);
796            let (mut expected, mut observed) = (0.0, 0.0);
797            for (v, p) in &d.outcomes {
798                expected += p * n as f64;
799                observed += seen.remove(&(v.as_f64().unwrap() as i64)).unwrap_or(0.0);
800                if expected >= 20.0 {
801                    chi2 += (observed - expected).powi(2) / expected;
802                    cells += 1;
803                    (expected, observed) = (0.0, 0.0);
804                }
805            }
806            chi2 += (observed - expected).powi(2) / expected.max(1.0);
807            assert!(seen.is_empty(), "{c:?} drew values outside its outcomes: {seen:?}");
808            let df = cells as f64;
809            assert!(
810                chi2 < df + 6.0 * (2.0 * df).sqrt() + 10.0,
811                "{c:?}: chi² {chi2} with {cells} cells"
812            );
813        }
814    }
815
816    #[test]
817    fn dice_pools() {
818        let b = &mut Budget::unlimited();
819        let pool = Dist::pool(3, &Dist::dice(1, 6, b).unwrap(), b).unwrap();
820        assert_eq!(pool.outcomes.len(), 56);
821        assert!(close(pool.total(), 1.0));
822        // A die with missing mass m: a pool of n misses 1 - (1 - m)^n.
823        let leaky = Dist::from_pairs(vec![(Value::Int(1.into()), 0.9)], 0.1);
824        let pool = Dist::pool(2, &leaky, b).unwrap();
825        assert!(close(pool.missing, 1.0 - 0.81));
826        assert!(close(pool.total() + pool.missing, 1.0));
827    }
828
829    #[test]
830    fn budgets_are_checked_before_building() {
831        let mut small = Budget {
832            max_outcomes: 1000,
833            max_collection: 1000,
834            ..Budget::unlimited()
835        };
836        assert!(Dist::dice(1, 100_000, &mut small).is_err());
837        assert!(Dist::pool(40, &Dist::dice(1, 6, &mut Budget::unlimited()).unwrap(), &mut small).is_err());
838        assert!(Dist::geometric(1e-9, &mut small).is_err());
839        assert_eq!(multisets(6, 3), 56);
840        assert_eq!(multisets(u64::MAX as u128, 40), u128::MAX);
841    }
842
843    #[test]
844    fn a_shared_budget_is_spent_once() {
845        let shared = Arc::new(AtomicU64::new(100_000));
846        let budget = Budget {
847            work_left: 0,
848            shared: Some(shared.clone()),
849            ..Budget::unlimited()
850        };
851        let (mut a, mut b) = (budget.clone(), budget);
852        a.work(60_000).unwrap();
853        assert!(b.work(60_000).is_err(), "only 40,000 are left");
854        b.work(30_000).unwrap();
855        a.give_back();
856        b.give_back();
857        assert_eq!(shared.load(Atomic::Relaxed), 10_000);
858    }
859
860    #[test]
861    fn gamma() {
862        assert!((ln_gamma(1.0)).abs() < 1e-13);
863        assert!((ln_gamma(5.0) - 24f64.ln()).abs() < 1e-13);
864        assert!((ln_gamma(0.5) - std::f64::consts::PI.sqrt().ln()).abs() < 1e-13);
865    }
866}