use super::super::{Cdf, Moments, Pdf, Quantile, Sample, bisection_quantile};
use crate::distributions::GammaDistribution;
use crate::rng::SplitMix64;
use crate::special::{gamma_p, ln_gamma};
impl Pdf for GammaDistribution {
fn pdf(&self, x: f64) -> f64 {
if x < 0.0 {
return 0.0;
}
let (k, theta) = (self.shape_parameter, self.scale_parameter);
let log_norm = k.mul_add(theta.ln(), ln_gamma(k));
let ln_density = (k - 1.0).mul_add(x.ln(), -(x / theta) - log_norm);
ln_density.exp()
}
}
impl Cdf for GammaDistribution {
fn cdf(&self, x: f64) -> f64 {
if x <= 0.0 {
0.0
} else {
gamma_p(self.shape_parameter, x / self.scale_parameter)
}
}
}
impl Quantile for GammaDistribution {
fn quantile(&self, p: f64) -> f64 {
let mean = self.shape_parameter * self.scale_parameter;
let sd = self.scale_parameter * self.shape_parameter.sqrt();
let hi = 20.0f64.mul_add(sd, mean).max(self.scale_parameter);
bisection_quantile(p, 0.0, hi, |x| self.cdf(x))
}
}
impl Moments for GammaDistribution {
fn mean(&self) -> Option<f64> {
Some(self.shape_parameter * self.scale_parameter)
}
fn variance(&self) -> Option<f64> {
Some(self.shape_parameter * self.scale_parameter * self.scale_parameter)
}
}
impl Sample for GammaDistribution {
fn sample(&self, rng: &mut SplitMix64) -> f64 {
self.scale_parameter * marsaglia_tsang(self.shape_parameter, rng)
}
}
pub(super) fn marsaglia_tsang(shape: f64, rng: &mut SplitMix64) -> f64 {
if shape < 1.0 {
let boosted = marsaglia_tsang(shape + 1.0, rng);
return boosted * rng.next_f64().powf(1.0 / shape);
}
let depth = shape - 1.0 / 3.0;
let coeff = 1.0 / (9.0 * depth).sqrt();
loop {
let normal = rng.standard_normal();
let base = coeff.mul_add(normal, 1.0);
if base <= 0.0 {
continue;
}
let cubed = base * base * base;
let uniform = rng.next_f64();
let normal_sq = normal * normal;
let accept = 0.5f64.mul_add(normal_sq, depth.mul_add(-cubed, depth) + depth * cubed.ln());
if uniform.ln() < accept {
return depth * cubed;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn shape_one_cdf_is_exponential() {
let d = GammaDistribution {
shape_parameter: 1.0,
scale_parameter: 2.0,
..Default::default()
};
assert!((d.cdf(2.0) - (1.0 - (-1.0f64).exp())).abs() < 1e-10);
}
#[test]
fn mean_is_shape_times_scale() {
let d = GammaDistribution {
shape_parameter: 2.5,
scale_parameter: 1.5,
..Default::default()
};
assert_eq!(d.mean(), Some(3.75));
}
}