use rand::rngs::StdRng;
use rand::{RngCore, SeedableRng};
pub trait RandomSource: RngCore {
fn reseed(&mut self, seed: u64) {
let _ = seed;
panic!("this random source does not support reseeding");
}
}
impl RandomSource for StdRng {
fn reseed(&mut self, seed: u64) {
*self = StdRng::seed_from_u64(seed);
}
}
#[derive(Debug, Clone)]
pub struct SplitMix64 {
state: u64,
}
impl SplitMix64 {
#[must_use]
pub fn new(seed: u64) -> Self {
SplitMix64 { state: seed }
}
pub fn set_seed(&mut self, seed: u64) {
self.state = seed;
}
#[inline]
fn next(&mut self) -> u64 {
self.state = self.state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
}
impl RngCore for SplitMix64 {
fn next_u64(&mut self) -> u64 {
self.next()
}
fn next_u32(&mut self) -> u32 {
(self.next() >> 32) as u32
}
fn fill_bytes(&mut self, dest: &mut [u8]) {
let mut chunks = dest.chunks_exact_mut(8);
for chunk in &mut chunks {
chunk.copy_from_slice(&self.next().to_le_bytes());
}
let rem = chunks.into_remainder();
if !rem.is_empty() {
let bytes = self.next().to_le_bytes();
rem.copy_from_slice(&bytes[..rem.len()]);
}
}
}
impl RandomSource for SplitMix64 {
fn reseed(&mut self, seed: u64) {
self.state = seed;
}
}
pub mod sample {
use rand::RngCore;
const TWO_POW_NEG_53: f64 = 1.0 / 9_007_199_254_740_992.0;
pub fn uniform01<R: RngCore + ?Sized>(rng: &mut R) -> f64 {
((rng.next_u64() >> 11) as f64) * TWO_POW_NEG_53
}
pub fn exponential<R: RngCore + ?Sized>(rng: &mut R, mean: f64) -> f64 {
let u = uniform01(rng);
-mean * (-u).ln_1p()
}
pub fn bernoulli<R: RngCore + ?Sized>(rng: &mut R, p: f64) -> bool {
uniform01(rng) < p
}
pub fn normal<R: RngCore + ?Sized>(rng: &mut R, mean: f64, std: f64) -> f64 {
let u1 = uniform01(rng);
let u2 = uniform01(rng);
let r = (-2.0 * (-u1).ln_1p()).sqrt();
let z = r * (std::f64::consts::TAU * u2).cos();
mean + std * z
}
}
#[cfg(test)]
mod tests {
use super::*;
const KAT_SEED0: [u64; 6] = [
0xE220_A839_7B1D_CDAF,
0x6E78_9E6A_A1B9_65F4,
0x06C4_5D18_8009_454F,
0xF88B_B8A8_724C_81EC,
0x1B39_896A_51A8_749B,
0x53CB_9F0C_747E_A2EA,
];
const KAT_SEED42: [u64; 6] = [
0xBDD7_3226_2FEB_6E95,
0x28EF_E333_B266_F103,
0x4752_6757_130F_9F52,
0x581C_E1FF_0E4A_E394,
0x09BC_585A_2448_23F2,
0xDE44_31FA_3C80_DB06,
];
#[test]
fn splitmix64_known_answer_vectors() {
for (seed, expected) in [(0u64, KAT_SEED0), (42u64, KAT_SEED42)] {
let mut rng = SplitMix64::new(seed);
for &want in &expected {
assert_eq!(rng.next_u64(), want, "seed {seed}");
}
}
}
#[test]
fn uniform01_and_exponential_known_answer() {
let mut rng = SplitMix64::new(0);
let u = sample::uniform01(&mut rng);
assert_eq!(u, 0.8833108082136426);
let mut rng = SplitMix64::new(0);
let e = sample::exponential(&mut rng, 1.0);
assert_eq!(e, 2.148241359348383);
}
#[test]
fn uniform01_in_unit_interval() {
let mut rng = SplitMix64::new(7);
for _ in 0..10_000 {
let u = sample::uniform01(&mut rng);
assert!((0.0..1.0).contains(&u));
}
}
#[test]
fn next_u32_is_high_bits_of_next_u64() {
let expected = (KAT_SEED0[0] >> 32) as u32;
let mut rng = SplitMix64::new(0);
assert_eq!(rng.next_u32(), expected);
}
#[test]
fn same_seed_same_stream() {
let mut a = SplitMix64::new(123);
let mut b = SplitMix64::new(123);
for _ in 0..1000 {
assert_eq!(a.next_u64(), b.next_u64());
}
}
#[test]
fn reseed_restarts_stream() {
let mut rng = SplitMix64::new(0);
let first = rng.next_u64();
let _ = rng.next_u64();
rng.reseed(0);
assert_eq!(rng.next_u64(), first);
rng.set_seed(0);
assert_eq!(rng.next_u64(), first);
}
#[test]
fn exponential_mean_is_sane() {
let mut rng = SplitMix64::new(99);
let n = 200_000;
let sum: f64 = (0..n).map(|_| sample::exponential(&mut rng, 5.0)).sum();
let mean = sum / f64::from(n);
assert!((mean - 5.0).abs() < 0.1, "mean was {mean}");
}
#[test]
fn normal_consumes_two_draws_and_is_centered() {
let mut a = SplitMix64::new(5);
let _ = sample::normal(&mut a, 0.0, 1.0);
let after_one = a.clone().next_u64();
let mut b = SplitMix64::new(5);
sample::uniform01(&mut b);
sample::uniform01(&mut b);
assert_eq!(after_one, b.next_u64());
let mut rng = SplitMix64::new(1);
let n = 200_000;
let sum: f64 = (0..n).map(|_| sample::normal(&mut rng, 10.0, 2.0)).sum();
let mean = sum / f64::from(n);
assert!((mean - 10.0).abs() < 0.05, "mean was {mean}");
}
#[test]
fn stdrng_reseed_is_deterministic() {
let mut a = StdRng::seed_from_u64(0);
a.reseed(77);
let mut b = StdRng::seed_from_u64(77);
assert_eq!(a.next_u64(), b.next_u64());
}
}