1use 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
13const TAIL: f64 = 1e-18;
15
16#[derive(Clone, Debug)]
17pub struct Dist {
18 pub outcomes: Vec<(Value, f64)>,
20 pub missing: f64,
24}
25
26const WORK_CHUNK: u64 = 1 << 14;
29
30#[derive(Clone, Debug)]
32pub struct Budget {
33 pub max_integer_bits: u64,
35 pub integer_bytes_left: Arc<AtomicU64>,
37 pub max_string_bytes: usize,
38 pub string_bytes_left: Arc<AtomicU64>,
40 pub cancel: Option<Arc<AtomicBool>>,
41 pub max_outcomes: usize,
43 pub max_collection: usize,
45 pub work_left: u64,
47 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 pub fn pool(count: u32, die: &Dist, budget: &mut Budget) -> OpResult<Dist> {
344 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 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 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 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 pub fn quantile(&self, q: f64) -> Option<Value> {
414 crate::stats::quantile(&self.outcomes, q).cloned()
415 }
416}
417
418#[derive(Clone, Copy, Debug)]
422pub enum Counts {
423 Binomial {
424 n: u64,
425 p: f64,
426 },
427 Poisson {
428 rate: f64,
429 },
430 Geometric {
432 p: f64,
433 },
434}
435
436impl Counts {
437 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 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 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 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 (libm::log(rng.open()) / libm::log1p(-p)).ceil().max(1.0) as i64
530 }
531 }
532 }
533}
534
535fn 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 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
574fn 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
627fn 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
635fn 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
697pub 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 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 #[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 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}