use super::super::{Cdf, LogCdf, Moments, Pdf, Quantile, Sample};
use crate::distributions::NormalDistribution;
use crate::rng::SplitMix64;
use crate::special::{erf, ln_erfc};
use std::f64::consts::{LN_2, PI, SQRT_2};
const INV_SQRT_2PI: f64 = 0.398_942_280_401_432_7;
impl NormalDistribution {
pub fn pdf_batch(&self, xs: &[f64], out: &mut [f64]) {
super::super::simd::normal_pdf_into(self, xs, out);
}
pub fn cdf_batch(&self, xs: &[f64], out: &mut [f64]) {
super::super::simd::normal_cdf_into(self, xs, out);
}
pub fn sample_batch(&self, rng: &mut SplitMix64, out: &mut [f64]) {
super::super::ziggurat::normal_sample_into(self.mean, self.standard_deviation, rng, out);
}
}
impl Pdf for NormalDistribution {
fn pdf(&self, x: f64) -> f64 {
let inv_sigma = 1.0 / self.standard_deviation;
let z = (x - self.mean) * inv_sigma;
(-0.5 * z * z).exp() * (INV_SQRT_2PI * inv_sigma)
}
}
impl Cdf for NormalDistribution {
fn cdf(&self, x: f64) -> f64 {
let inv_scale = 1.0 / (self.standard_deviation * SQRT_2);
let z = (x - self.mean) * inv_scale;
0.5 * (1.0 + erf(z))
}
}
impl LogCdf for NormalDistribution {
fn logsf(&self, x: f64) -> f64 {
let z = (x - self.mean) / (self.standard_deviation * SQRT_2);
ln_erfc(z) - LN_2
}
fn logcdf(&self, x: f64) -> f64 {
let z = (x - self.mean) / (self.standard_deviation * SQRT_2);
ln_erfc(-z) - LN_2
}
}
impl Quantile for NormalDistribution {
fn quantile(&self, p: f64) -> f64 {
let erf_target = 2.0f64.mul_add(p, -1.0);
self.standard_deviation
.mul_add(SQRT_2 * inv_erf(erf_target), self.mean)
}
}
impl Moments for NormalDistribution {
fn mean(&self) -> Option<f64> {
Some(self.mean)
}
fn variance(&self) -> Option<f64> {
Some(self.standard_deviation * self.standard_deviation)
}
}
impl Sample for NormalDistribution {
fn sample(&self, rng: &mut SplitMix64) -> f64 {
self.standard_deviation
.mul_add(rng.standard_normal(), self.mean)
}
}
fn inv_erf(y: f64) -> f64 {
if y <= -1.0 {
return f64::NEG_INFINITY;
}
if y >= 1.0 {
return f64::INFINITY;
}
let w = -((1.0 - y) * (1.0 + y)).ln();
let mut x = if w < 5.0 {
let w = w - 2.5;
let mut p: f64 = 2.810_226_36e-08;
for c in [
3.432_739_39e-07,
-3.523_387_7e-06,
-4.391_506_54e-06,
0.000_218_580_87,
-0.001_253_725_03,
-0.004_177_681_64,
0.246_640_727,
1.501_409_41,
] {
p = p.mul_add(w, c);
}
p * y
} else {
let w = w.sqrt() - 3.0;
let mut p: f64 = -0.000_200_214_257;
for c in [
0.000_100_950_558,
0.001_349_343_22,
-0.003_673_428_44,
0.005_739_507_73,
-0.007_622_461_3,
0.009_438_870_47,
1.001_674_06,
2.832_976_82,
] {
p = p.mul_add(w, c);
}
p * y
};
let two_over_sqrt_pi = 2.0 / PI.sqrt();
for _ in 0..3 {
let err = erf(x) - y;
x -= err / (two_over_sqrt_pi * (-x * x).exp());
}
x
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn standard_normal_peak_density() {
let n = NormalDistribution {
mean: 0.0,
standard_deviation: 1.0,
..Default::default()
};
assert!(
(n.pdf(0.0) - 0.398_942_28).abs() < 1e-6,
"peak density was {}",
n.pdf(0.0)
);
}
#[test]
fn density_decreases_away_from_mean() {
fn density_at(d: &impl Pdf, x: f64) -> f64 {
d.pdf(x)
}
let n = NormalDistribution {
mean: 2.0,
standard_deviation: 0.5,
..Default::default()
};
assert!(
density_at(&n, 2.0) > density_at(&n, 3.0),
"density should fall off from the mean"
);
}
#[test]
fn logsf_matches_scipy_and_stays_finite() {
let n = NormalDistribution {
mean: 0.0,
standard_deviation: 1.0,
..Default::default()
};
let body = n.logsf(0.5);
assert!(
((body - (1.0 - n.cdf(0.5)).ln()) / body.abs()).abs() < 1e-10,
"logsf body was {body}"
);
let tail = n.logsf(40.0);
let want = -804.608_442_013_753_9;
assert!(tail.is_finite(), "logsf(40) was {tail}");
assert!(
((tail - want) / want).abs() < 1e-9,
"logsf(40) = {tail}, want {want}"
);
}
}