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