#![cfg(feature = "alloc")]
#![cfg(feature = "std")]
use super::traits::Distribution;
use rand::Rng;
use rand::RngExt;
pub(crate) fn sample_standard_normal<R: Rng + ?Sized>(rng: &mut R) -> f64 {
loop {
let u = 2.0 * rng.random::<f64>() - 1.0;
let v = 2.0 * rng.random::<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 sample_exp1<R: Rng + ?Sized>(rng: &mut R) -> f64 {
-(1.0 - rng.random::<f64>()).ln()
}
#[derive(Clone, Copy, Debug)]
pub struct Normal {
mean: f64,
std_dev: f64,
}
impl Normal {
pub fn new(mean: f64, std_dev: f64) -> Self {
assert!(
std_dev.is_finite() && std_dev >= 0.0,
"Normal std_dev must be finite and non-negative"
);
Self { mean, std_dev }
}
}
impl Distribution for Normal {
type Output = f64;
#[inline]
fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> f64 {
self.mean + self.std_dev * sample_standard_normal(rng)
}
}
#[derive(Clone, Copy, Debug)]
pub struct Uniform {
dist: rand::distr::Uniform<f64>,
}
impl Uniform {
pub fn new(min: f64, max: f64) -> Self {
Self {
dist: rand::distr::Uniform::new(min, max).unwrap(),
}
}
}
impl Distribution for Uniform {
type Output = f64;
#[inline]
fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> f64 {
rng.sample(self.dist)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::SeedableRng;
use rand::rngs::StdRng;
#[test]
fn normal_samples_are_finite() {
let dist = Normal::new(0.0, 1.0);
let mut rng = StdRng::seed_from_u64(0xdead_beef);
for _ in 0..256 {
assert!(
dist.sample(&mut rng).is_finite(),
"Normal(0, 1) sample must be a finite f64"
);
}
}
#[test]
#[should_panic(expected = "EmptyRange")]
fn uniform_panics_on_inverted_bounds() {
let _ = Uniform::new(1.0, 0.0);
}
#[test]
#[allow(clippy::float_cmp)] fn normal_zero_std_dev_always_returns_mean() {
let mean = std::f64::consts::PI;
let dist = Normal::new(mean, 0.0);
let mut rng = StdRng::seed_from_u64(0x1234_5678_9abc_def0);
for _ in 0..64 {
let sample = dist.sample(&mut rng);
assert_eq!(
sample, mean,
"Normal(mean={mean}, std_dev=0) must always sample exactly the mean"
);
}
}
#[test]
fn normal_moments_match() {
let dist = Normal::new(2.0, 3.0);
let mut rng = StdRng::seed_from_u64(42);
let n = 100_000_i32;
let n_f = f64::from(n);
let samples: Vec<f64> = (0..n).map(|_| dist.sample(&mut rng)).collect();
let mean = samples.iter().sum::<f64>() / n_f;
let var = samples.iter().map(|x| (x - mean) * (x - mean)).sum::<f64>() / n_f;
assert!((mean - 2.0).abs() < 0.05, "mean {mean} too far from 2.0");
assert!((var - 9.0).abs() < 0.3, "variance {var} too far from 9.0");
}
#[test]
fn exp1_moments_match() {
let mut rng = StdRng::seed_from_u64(43);
let n = 100_000_i32;
let n_f = f64::from(n);
let samples: Vec<f64> = (0..n).map(|_| sample_exp1(&mut rng)).collect();
for &s in &samples {
assert!(
s.is_finite() && s >= 0.0,
"Exp(1) sample {s} out of support"
);
}
let mean = samples.iter().sum::<f64>() / n_f;
let var = samples.iter().map(|x| (x - mean) * (x - mean)).sum::<f64>() / n_f;
assert!((mean - 1.0).abs() < 0.05, "mean {mean} too far from 1.0");
assert!((var - 1.0).abs() < 0.1, "variance {var} too far from 1.0");
}
#[test]
fn uniform_samples_are_in_range() {
let min = 2.5_f64;
let max = 7.5_f64;
let dist = Uniform::new(min, max);
let mut rng = StdRng::seed_from_u64(0xcafe_babe);
for _ in 0..256 {
let sample = dist.sample(&mut rng);
assert!(
sample >= min && sample < max,
"Uniform({min}, {max}) sample {sample} is out of range"
);
}
}
}