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 #[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 #[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 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 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 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 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 #[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 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 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 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 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 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 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 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 pub fn pool(count: u32, die: &Dist, budget: &mut Budget) -> OpResult<Dist> {
350 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 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 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 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 pub fn quantile(&self, q: f64) -> Option<Value> {
420 crate::stats::quantile(&self.outcomes, q).cloned()
421 }
422}
423
424#[derive(Clone, Copy, Debug)]
428pub enum Counts {
429 Binomial {
430 n: u64,
431 p: f64,
432 },
433 Poisson {
434 rate: f64,
435 },
436 Geometric {
438 p: f64,
439 },
440}
441
442impl Counts {
443 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 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 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 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 (libm::log(rng.open()) / libm::log1p(-p)).ceil().max(1.0) as i64
536 }
537 }
538 }
539}
540
541fn 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 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
580fn 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
633fn 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
641fn 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
703pub 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 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 #[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 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}