1use prospicio_core::{Error, Result, StreamRng};
8use prospicio_math::special::ln_gamma;
9
10use crate::distribution::check_probability;
11
12pub trait Counting {
14 fn pmf(&self, k: u64) -> f64;
16
17 fn mean(&self) -> f64;
19
20 fn variance(&self) -> f64;
22
23 fn panjer_ab(&self) -> (f64, f64);
25
26 fn pgf(&self, z: f64) -> f64;
29
30 fn pgf_complex(&self, z: (f64, f64)) -> (f64, f64);
33
34 fn cdf(&self, k: u64) -> f64 {
36 (0..=k).map(|j| self.pmf(j)).sum::<f64>().min(1.0)
37 }
38
39 fn quantile(&self, p: f64) -> Result<u64> {
41 check_probability(p)?;
42 if p == 1.0 {
43 return Ok(u64::MAX);
44 }
45 let mut total = 0.0;
46 let mut k = 0;
47 loop {
48 let q = self.pmf(k);
49 total += q;
50 if total >= p || (q == 0.0 && k as f64 > self.mean()) {
52 return Ok(k);
53 }
54 k += 1;
55 }
56 }
57
58 fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<u64> {
64 (0..n)
65 .map(|_| {
66 self.quantile(rng.next_open01())
67 .expect("next_open01 is always in (0, 1)")
68 })
69 .collect()
70 }
71}
72
73impl<T: Counting + ?Sized> Counting for Box<T> {
74 fn pmf(&self, k: u64) -> f64 {
75 (**self).pmf(k)
76 }
77
78 fn mean(&self) -> f64 {
79 (**self).mean()
80 }
81
82 fn variance(&self) -> f64 {
83 (**self).variance()
84 }
85
86 fn panjer_ab(&self) -> (f64, f64) {
87 (**self).panjer_ab()
88 }
89
90 fn pgf(&self, z: f64) -> f64 {
91 (**self).pgf(z)
92 }
93
94 fn pgf_complex(&self, z: (f64, f64)) -> (f64, f64) {
95 (**self).pgf_complex(z)
96 }
97
98 fn cdf(&self, k: u64) -> f64 {
99 (**self).cdf(k)
100 }
101
102 fn quantile(&self, p: f64) -> Result<u64> {
103 (**self).quantile(p)
104 }
105
106 fn sample(&self, rng: &mut StreamRng, n: usize) -> Vec<u64> {
107 (**self).sample(rng, n)
108 }
109}
110
111#[derive(Debug, Clone, Copy, PartialEq)]
121pub struct Poisson {
122 lambda: f64,
123}
124
125impl Poisson {
126 pub fn new(lambda: f64) -> Result<Self> {
129 if !lambda.is_finite() || lambda < 0.0 {
130 return Err(Error::InvalidParameter {
131 name: "lambda",
132 value: lambda,
133 reason: "must be finite and non-negative",
134 });
135 }
136 Ok(Self { lambda })
137 }
138
139 pub fn lambda(&self) -> f64 {
141 self.lambda
142 }
143}
144
145impl Counting for Poisson {
146 fn pmf(&self, k: u64) -> f64 {
147 if self.lambda == 0.0 {
148 return if k == 0 { 1.0 } else { 0.0 };
149 }
150 let k = k as f64;
151 (k * self.lambda.ln() - self.lambda - ln_gamma(k + 1.0)).exp()
152 }
153
154 fn mean(&self) -> f64 {
155 self.lambda
156 }
157
158 fn variance(&self) -> f64 {
159 self.lambda
160 }
161
162 fn panjer_ab(&self) -> (f64, f64) {
163 (0.0, self.lambda)
164 }
165
166 fn pgf(&self, z: f64) -> f64 {
168 (self.lambda * (z - 1.0)).exp()
169 }
170
171 fn pgf_complex(&self, (re, im): (f64, f64)) -> (f64, f64) {
172 let modulus = (self.lambda * (re - 1.0)).exp();
173 let angle = self.lambda * im;
174 (modulus * angle.cos(), modulus * angle.sin())
175 }
176}
177
178#[derive(Debug, Clone, Copy, PartialEq)]
192pub struct NegativeBinomial {
193 r: f64,
194 beta: f64,
195}
196
197impl NegativeBinomial {
198 pub fn new(r: f64, beta: f64) -> Result<Self> {
201 if !r.is_finite() || r <= 0.0 {
202 return Err(Error::InvalidParameter {
203 name: "r",
204 value: r,
205 reason: "must be finite and positive",
206 });
207 }
208 if !beta.is_finite() || beta <= 0.0 {
209 return Err(Error::InvalidParameter {
210 name: "beta",
211 value: beta,
212 reason: "must be finite and positive",
213 });
214 }
215 Ok(Self { r, beta })
216 }
217
218 pub fn from_mean_variance(mean: f64, variance: f64) -> Result<Self> {
221 if !mean.is_finite() || mean <= 0.0 {
222 return Err(Error::InvalidParameter {
223 name: "mean",
224 value: mean,
225 reason: "must be finite and positive",
226 });
227 }
228 if !variance.is_finite() || variance <= mean {
229 return Err(Error::InvalidParameter {
230 name: "variance",
231 value: variance,
232 reason: "must exceed the mean",
233 });
234 }
235 let beta = variance / mean - 1.0;
236 Self::new(mean / beta, beta)
237 }
238
239 pub fn r(&self) -> f64 {
241 self.r
242 }
243
244 pub fn beta(&self) -> f64 {
246 self.beta
247 }
248}
249
250impl Counting for NegativeBinomial {
251 fn pmf(&self, k: u64) -> f64 {
252 let (r, beta) = (self.r, self.beta);
253 let k = k as f64;
254 let ln_1p_beta = beta.ln_1p();
255 (ln_gamma(k + r) - ln_gamma(r) - ln_gamma(k + 1.0) - r * ln_1p_beta
256 + k * (beta.ln() - ln_1p_beta))
257 .exp()
258 }
259
260 fn mean(&self) -> f64 {
261 self.r * self.beta
262 }
263
264 fn variance(&self) -> f64 {
265 self.r * self.beta * (1.0 + self.beta)
266 }
267
268 fn panjer_ab(&self) -> (f64, f64) {
269 let a = self.beta / (1.0 + self.beta);
270 (a, (self.r - 1.0) * a)
271 }
272
273 fn pgf(&self, z: f64) -> f64 {
275 (-self.r * (self.beta * (1.0 - z)).ln_1p()).exp()
276 }
277
278 fn pgf_complex(&self, (re, im): (f64, f64)) -> (f64, f64) {
279 let (wr, wi) = (1.0 + self.beta * (1.0 - re), -self.beta * im);
281 let ln_modulus = 0.5 * (wr * wr + wi * wi).ln();
282 let arg = wi.atan2(wr);
283 let modulus = (-self.r * ln_modulus).exp();
284 let angle = -self.r * arg;
285 (modulus * angle.cos(), modulus * angle.sin())
286 }
287}
288
289#[derive(Debug, Clone, Copy, PartialEq)]
300pub struct Binomial {
301 n: u64,
302 p: f64,
303}
304
305impl Binomial {
306 pub fn new(n: u64, p: f64) -> Result<Self> {
309 if !(0.0..1.0).contains(&p) {
310 return Err(Error::InvalidParameter {
311 name: "p",
312 value: p,
313 reason: "must be in [0, 1)",
314 });
315 }
316 Ok(Self { n, p })
317 }
318
319 pub fn n(&self) -> u64 {
321 self.n
322 }
323
324 pub fn p(&self) -> f64 {
326 self.p
327 }
328}
329
330impl Counting for Binomial {
331 fn pmf(&self, k: u64) -> f64 {
332 if k > self.n {
333 return 0.0;
334 }
335 if self.p == 0.0 {
336 return if k == 0 { 1.0 } else { 0.0 };
337 }
338 let (n, k) = (self.n as f64, k as f64);
339 (ln_gamma(n + 1.0) - ln_gamma(k + 1.0) - ln_gamma(n - k + 1.0)
340 + k * self.p.ln()
341 + (n - k) * (-self.p).ln_1p())
342 .exp()
343 }
344
345 fn mean(&self) -> f64 {
346 self.n as f64 * self.p
347 }
348
349 fn variance(&self) -> f64 {
350 self.n as f64 * self.p * (1.0 - self.p)
351 }
352
353 fn panjer_ab(&self) -> (f64, f64) {
354 let odds = self.p / (1.0 - self.p);
355 (-odds, (self.n as f64 + 1.0) * odds)
356 }
357
358 fn pgf(&self, z: f64) -> f64 {
360 (1.0 + self.p * (z - 1.0)).powf(self.n as f64)
361 }
362
363 fn pgf_complex(&self, (re, im): (f64, f64)) -> (f64, f64) {
364 let (wr, wi) = (1.0 + self.p * (re - 1.0), self.p * im);
366 let n = self.n as f64;
367 let modulus = (0.5 * n * (wr * wr + wi * wi).ln()).exp();
368 let angle = n * wi.atan2(wr);
369 (modulus * angle.cos(), modulus * angle.sin())
370 }
371}
372
373#[derive(Debug, Clone, Copy, PartialEq)]
386pub enum PanjerClass {
387 Binomial(Binomial),
388 Poisson(Poisson),
389 NegativeBinomial(NegativeBinomial),
390}
391
392impl PanjerClass {
393 pub fn from_mean_dispersion(mean: f64, dispersion: f64) -> Result<Self> {
402 if !mean.is_finite() || mean < 0.0 {
403 return Err(Error::InvalidParameter {
404 name: "mean",
405 value: mean,
406 reason: "must be finite and non-negative",
407 });
408 }
409 if !dispersion.is_finite() || dispersion <= 0.0 {
410 return Err(Error::InvalidParameter {
411 name: "dispersion",
412 value: dispersion,
413 reason: "must be finite and positive",
414 });
415 }
416 if dispersion == 1.0 || mean == 0.0 {
417 return Ok(Self::Poisson(Poisson::new(mean)?));
418 }
419 if dispersion > 1.0 {
420 let beta = dispersion - 1.0;
421 return Ok(Self::NegativeBinomial(NegativeBinomial::new(
422 mean / beta,
423 beta,
424 )?));
425 }
426 let trials = mean / (1.0 - dispersion);
427 let n = if (trials - trials.round()).abs() <= 1e-9 * trials {
429 trials.round()
430 } else {
431 trials.ceil()
432 };
433 Ok(Self::Binomial(Binomial::new(
434 n as u64,
435 (mean / n).min(1.0),
436 )?))
437 }
438
439 pub fn dispersion(&self) -> f64 {
441 match self {
442 Self::Binomial(b) => 1.0 - b.p(),
443 Self::Poisson(_) => 1.0,
444 Self::NegativeBinomial(nb) => 1.0 + nb.beta(),
445 }
446 }
447
448 fn inner(&self) -> &dyn Counting {
449 match self {
450 Self::Binomial(n) => n,
451 Self::Poisson(n) => n,
452 Self::NegativeBinomial(n) => n,
453 }
454 }
455}
456
457impl Counting for PanjerClass {
458 fn pmf(&self, k: u64) -> f64 {
459 self.inner().pmf(k)
460 }
461
462 fn mean(&self) -> f64 {
463 self.inner().mean()
464 }
465
466 fn variance(&self) -> f64 {
467 self.inner().variance()
468 }
469
470 fn panjer_ab(&self) -> (f64, f64) {
471 self.inner().panjer_ab()
472 }
473
474 fn pgf(&self, z: f64) -> f64 {
475 self.inner().pgf(z)
476 }
477
478 fn pgf_complex(&self, z: (f64, f64)) -> (f64, f64) {
479 self.inner().pgf_complex(z)
480 }
481}
482
483#[cfg(test)]
484mod tests {
485 use super::*;
486
487 fn check_recursion(n: &impl Counting) {
489 let (a, b) = n.panjer_ab();
490 for k in 1..60u64 {
491 let want = (a + b / k as f64) * n.pmf(k - 1);
492 let got = n.pmf(k);
493 assert!((got - want).abs() <= 1e-12 * want.max(1e-300), "k {k}");
494 }
495 }
496
497 #[test]
498 fn panjer_recursion_holds() {
499 check_recursion(&Poisson::new(4.5).unwrap());
500 check_recursion(&NegativeBinomial::new(2.5, 1.5).unwrap());
501 check_recursion(&NegativeBinomial::new(0.7, 10.0).unwrap());
502 check_recursion(&Binomial::new(12, 0.35).unwrap());
503 check_recursion(&Binomial::new(300, 0.02).unwrap());
504 }
505
506 #[test]
507 fn panjer_class_by_dispersion() {
508 let nb = PanjerClass::from_mean_dispersion(4.0, 3.0).unwrap();
509 assert!(matches!(nb, PanjerClass::NegativeBinomial(_)));
510 assert!((nb.variance() / nb.mean() - 3.0).abs() < 1e-14);
511 let po = PanjerClass::from_mean_dispersion(4.0, 1.0).unwrap();
512 assert_eq!(po, PanjerClass::Poisson(Poisson::new(4.0).unwrap()));
513 let bi = PanjerClass::from_mean_dispersion(10.0, 0.5).unwrap();
515 assert_eq!(bi, PanjerClass::Binomial(Binomial::new(20, 0.5).unwrap()));
516 let up = PanjerClass::from_mean_dispersion(10.0, 0.3).unwrap();
518 let PanjerClass::Binomial(b) = up else {
519 panic!("{up:?}")
520 };
521 assert_eq!(b.n(), 15);
522 assert!((up.mean() - 10.0).abs() < 1e-14);
523 assert!((up.dispersion() - 1.0 / 3.0).abs() < 1e-15);
524 assert!((up.variance() / up.mean() - up.dispersion()).abs() < 1e-15);
525 assert_eq!(
526 PanjerClass::from_mean_dispersion(0.0, 2.0).unwrap(),
527 PanjerClass::Poisson(Poisson::new(0.0).unwrap())
528 );
529 assert!(PanjerClass::from_mean_dispersion(1.0, 0.0).is_err());
530 assert!(PanjerClass::from_mean_dispersion(-1.0, 1.0).is_err());
531 check_recursion(&up);
532 }
533
534 #[test]
535 fn pmf_sums_to_one_and_matches_moments() {
536 for n in [
537 &Poisson::new(7.0).unwrap() as &dyn Counting,
538 &NegativeBinomial::new(3.0, 2.0).unwrap(),
539 &Binomial::new(40, 0.15).unwrap(),
540 ] {
541 let pmf: Vec<f64> = (0..400).map(|k| n.pmf(k)).collect();
542 let total: f64 = pmf.iter().sum();
543 let mean: f64 = pmf.iter().enumerate().map(|(k, p)| k as f64 * p).sum();
544 let var: f64 = pmf
545 .iter()
546 .enumerate()
547 .map(|(k, p)| (k as f64 - mean).powi(2) * p)
548 .sum();
549 assert!((total - 1.0).abs() < 1e-12);
550 assert!((mean - n.mean()).abs() < 1e-10);
551 assert!((var - n.variance()).abs() < 1e-9);
552 }
553 }
554
555 #[test]
556 fn quantile_and_cdf_are_inverse() {
557 let n = Poisson::new(3.0).unwrap();
558 for k in 0..15 {
559 assert_eq!(n.quantile(n.cdf(k)), Ok(k));
560 }
561 assert_eq!(n.quantile(0.0), Ok(0));
562 assert_eq!(n.quantile(1.0), Ok(u64::MAX));
563 assert!(n.quantile(1.5).is_err());
564 }
565
566 #[test]
567 fn sampling_is_reproducible_and_centred() {
568 let n = NegativeBinomial::new(4.0, 2.5).unwrap();
569 let a = n.sample(&mut StreamRng::new(5, 0), 100_000);
570 assert_eq!(a, n.sample(&mut StreamRng::new(5, 0), 100_000));
571 let mean = a.iter().sum::<u64>() as f64 / a.len() as f64;
572 assert!((mean - 10.0).abs() < 0.1, "{mean}");
574 }
575
576 #[test]
577 fn pgf_matches_the_pmf() {
578 for n in [
579 &Poisson::new(2.5).unwrap() as &dyn Counting,
580 &NegativeBinomial::new(1.5, 3.0).unwrap(),
581 &Binomial::new(25, 0.2).unwrap(),
582 ] {
583 for z in [0.0f64, 0.3, 0.9, 1.0] {
584 let series: f64 = (0..500).map(|k| n.pmf(k) * z.powi(k as i32)).sum();
585 assert!((n.pgf(z) - series).abs() < 1e-12, "z {z}");
586 }
587 }
588 }
589
590 #[test]
591 fn complex_pgf_matches_the_series() {
592 for n in [
593 &Poisson::new(2.5).unwrap() as &dyn Counting,
594 &NegativeBinomial::new(1.5, 3.0).unwrap(),
595 &Binomial::new(25, 0.2).unwrap(),
596 ] {
597 for theta in [0.0f64, 0.7, 2.0, 3.1] {
598 let (zr, zi) = (0.9 * theta.cos(), 0.9 * theta.sin());
599 let (mut pr, mut pi, mut sr, mut si) = (1.0, 0.0, 0.0, 0.0);
601 for k in 0..600 {
602 let p = n.pmf(k);
603 sr += p * pr;
604 si += p * pi;
605 (pr, pi) = (pr * zr - pi * zi, pr * zi + pi * zr);
606 }
607 let (gr, gi) = n.pgf_complex((zr, zi));
608 assert!(
609 (gr - sr).abs() < 1e-12 && (gi - si).abs() < 1e-12,
610 "theta {theta}"
611 );
612 }
613 assert!((n.pgf_complex((0.4, 0.0)).0 - n.pgf(0.4)).abs() < 1e-15);
614 }
615 }
616
617 #[test]
618 fn degenerate_poisson() {
619 let n = Poisson::new(0.0).unwrap();
620 assert_eq!(n.pmf(0), 1.0);
621 assert_eq!(n.pmf(3), 0.0);
622 assert_eq!(n.quantile(0.999), Ok(0));
623 }
624
625 #[test]
626 fn rejects_bad_parameters() {
627 assert!(Poisson::new(-1.0).is_err());
628 assert!(NegativeBinomial::new(0.0, 1.0).is_err());
629 assert!(NegativeBinomial::new(1.0, f64::NAN).is_err());
630 assert!(NegativeBinomial::from_mean_variance(10.0, 10.0).is_err());
631 assert!(Binomial::new(5, 1.0).is_err());
632 assert!(Binomial::new(5, -0.1).is_err());
633 let zero = Binomial::new(5, 0.0).unwrap();
634 assert_eq!((zero.pmf(0), zero.pmf(1)), (1.0, 0.0));
635 }
636}