use crate::output::codegen::numeric_rt::lgamma;
pub use crate::base::rng::SplitMix64 as Rng;
pub(crate) fn positive_uniform(rng: &mut Rng) -> f64 {
1.0 - rng.next_f64()
}
pub(crate) fn standard_normal(rng: &mut Rng) -> f64 {
loop {
let u = 2.0 * rng.next_f64() - 1.0;
let v = 2.0 * rng.next_f64() - 1.0;
let s = u * u + v * v;
if s > 0.0 && s < 1.0 {
return u * (-2.0 * s.ln() / s).sqrt();
}
}
}
pub(crate) fn standard_gamma(rng: &mut Rng, shape: f64) -> f64 {
if shape < 1.0 {
let boost = positive_uniform(rng).powf(1.0 / shape);
return standard_gamma(rng, shape + 1.0) * boost;
}
let d = shape - 1.0 / 3.0;
let c = 1.0 / (9.0 * d).sqrt();
loop {
let x = standard_normal(rng);
let t = 1.0 + c * x;
if t <= 0.0 {
continue;
}
let v = t * t * t;
let u = rng.next_f64();
let x2 = x * x;
if u < 1.0 - 0.0331 * x2 * x2 {
return d * v;
}
if u.ln() < 0.5 * x2 + d * (1.0 - v + v.ln()) {
return d * v;
}
}
}
const POISSON_KNUTH_LIMIT: f64 = 30.0;
pub(crate) fn poisson(rng: &mut Rng, lambda: f64) -> f64 {
if lambda < POISSON_KNUTH_LIMIT {
let limit = (-lambda).exp();
let mut k = 0.0;
let mut product = 1.0;
loop {
product *= rng.next_f64();
if product <= limit {
return k;
}
k += 1.0;
}
}
let log_lambda = lambda.ln();
let b = 0.931 + 2.53 * lambda.sqrt();
let a = -0.059 + 0.02483 * b;
let inv_alpha = 1.1239 + 1.1328 / (b - 3.4);
let v_r = 0.9277 - 3.6224 / (b - 2.0);
loop {
let u = rng.next_f64() - 0.5;
let v = rng.next_f64();
let us = 0.5 - u.abs();
let k = ((2.0 * a / us + b) * u + lambda + 0.43).floor();
if us >= 0.07 && v <= v_r {
return k;
}
if k < 0.0 || (us < 0.013 && v > us) {
continue;
}
if v.ln() + inv_alpha.ln() - (a / (us * us) + b).ln()
<= -lambda + k * log_lambda - lgamma(k + 1.0)
{
return k;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn moments(samples: &[f64]) -> (f64, f64) {
let n = samples.len() as f64;
let mean = samples.iter().sum::<f64>() / n;
let var = samples.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (n - 1.0);
(mean, var)
}
#[test]
fn standard_normal_moments() {
let mut rng = Rng::new(3);
let s: Vec<f64> = (0..20_000).map(|_| standard_normal(&mut rng)).collect();
let (mean, var) = moments(&s);
assert!(mean.abs() < 4.0 / 20_000f64.sqrt(), "mean {mean}");
assert!((var - 1.0).abs() < 0.05, "variance {var}");
}
#[test]
fn standard_gamma_both_branches() {
for shape in [0.3, 0.9, 1.0, 2.5, 17.0] {
let mut rng = Rng::new(5);
let s: Vec<f64> = (0..20_000)
.map(|_| standard_gamma(&mut rng, shape))
.collect();
assert!(s.iter().all(|v| *v >= 0.0), "shape {shape}: negative draw");
let (mean, var) = moments(&s);
let se = shape.sqrt() / 20_000f64.sqrt();
assert!(
(mean - shape).abs() < 4.0 * se,
"shape {shape}: mean {mean}"
);
assert!(
(var - shape).abs() < 0.1 * shape.max(1.0),
"shape {shape}: variance {var}"
);
}
}
#[test]
fn poisson_both_branches() {
for lambda in [0.5, 4.0, 29.9, 30.0, 150.0] {
let mut rng = Rng::new(7);
let s: Vec<f64> = (0..20_000).map(|_| poisson(&mut rng, lambda)).collect();
assert!(
s.iter().all(|v| *v >= 0.0 && v.fract() == 0.0),
"λ = {lambda}: non-lattice draw"
);
let (mean, var) = moments(&s);
let se = lambda.sqrt() / 20_000f64.sqrt();
assert!(
(mean - lambda).abs() < 4.0 * se,
"λ = {lambda}: mean {mean}"
);
assert!(
(var - lambda).abs() < 0.1 * lambda.max(1.0),
"λ = {lambda}: variance {var}"
);
}
}
}