use std::ops::Range;
pub(crate) const GOLDEN: u64 = 0x9E37_79B9_7F4A_7C15;
#[inline]
pub(crate) fn mix64(mut z: u64) -> u64 {
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
#[inline]
pub(crate) fn splitmix64(z: u64) -> u64 {
mix64(z.wrapping_add(GOLDEN))
}
pub(crate) fn stream_key(parts: &[u64]) -> u64 {
parts.iter().fold(0, |key, &part| splitmix64(key ^ part))
}
#[inline]
fn keyed_bits(key: u64, index: u64) -> u64 {
mix64(key.wrapping_add(index.wrapping_add(1).wrapping_mul(GOLDEN)))
}
#[inline]
fn keyed_unit_open(key: u64, index: u64) -> f64 {
((keyed_bits(key, index) >> 11) + 1) as f64 * (1.0 / (1u64 << 53) as f64)
}
#[inline]
pub(crate) fn keyed_unit(key: u64, index: u64) -> f64 {
(keyed_bits(key, index) >> 11) as f64 * (1.0 / (1u64 << 53) as f64)
}
#[inline]
pub(crate) fn keyed_unit_f32(key: u64, index: u64) -> f32 {
(keyed_bits(key, index) >> 40) as f32 * (1.0 / (1u32 << 24) as f32)
}
#[inline]
pub(crate) fn keyed_normal(key: u64, index: u64) -> f64 {
let u1 = keyed_unit_open(key, index.wrapping_mul(2));
let u2 = keyed_unit_open(key, index.wrapping_mul(2).wrapping_add(1));
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
#[derive(Debug, Clone)]
pub(crate) struct Rng {
s: [u64; 4],
}
impl Rng {
pub(crate) fn new(seed: u64) -> Self {
let mut z = seed;
let mut s = [0; 4];
for word in &mut s {
z = z.wrapping_add(GOLDEN);
*word = mix64(z);
}
Rng { s }
}
#[inline]
pub(crate) fn next_u64(&mut self) -> u64 {
let s = &mut self.s;
let out = s[0].wrapping_add(s[3]).rotate_left(23).wrapping_add(s[0]);
let t = s[1] << 17;
s[2] ^= s[0];
s[3] ^= s[1];
s[1] ^= s[2];
s[0] ^= s[3];
s[2] ^= t;
s[3] = s[3].rotate_left(45);
out
}
#[inline]
pub(crate) fn f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 * (1.0 / (1u32 << 24) as f32)
}
#[inline]
pub(crate) fn f64(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 * (1.0 / (1u64 << 53) as f64)
}
#[inline]
pub(crate) fn below(&mut self, n: u64) -> u64 {
debug_assert!(n > 0, "empty sampling range");
let mut m = u128::from(self.next_u64()) * u128::from(n);
if (m as u64) < n {
let threshold = n.wrapping_neg() % n;
while (m as u64) < threshold {
m = u128::from(self.next_u64()) * u128::from(n);
}
}
(m >> 64) as u64
}
#[inline]
pub(crate) fn range(&mut self, range: Range<usize>) -> usize {
debug_assert!(range.start < range.end, "empty sampling range");
range.start + self.below((range.end - range.start) as u64) as usize
}
pub(crate) fn shuffle<T>(&mut self, items: &mut [T]) {
for i in (1..items.len()).rev() {
items.swap(i, self.below(i as u64 + 1) as usize);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn matches_the_reference_stream() {
let mut rng = Rng { s: [1, 2, 3, 4] };
let expected = [
41_943_041,
58_720_359,
3_588_806_011_781_223,
3_591_011_842_654_386,
9_228_616_714_210_784_205,
];
for e in expected {
assert_eq!(rng.next_u64(), e);
}
}
#[test]
fn bounded_draws_are_unbiased_and_in_range() {
let mut rng = Rng::new(7);
let n = 6;
let mut counts = [0u32; 6];
for _ in 0..60_000 {
counts[rng.below(n) as usize] += 1;
}
assert!(
counts.iter().all(|&c| c.abs_diff(10_000) < 500),
"{counts:?}"
);
let big = (1u64 << 63) + 1;
assert!((0..1000).all(|_| rng.below(big) < big));
assert!((0..1000).all(|_| (5..8).contains(&rng.range(5..8))));
}
#[test]
fn shuffle_is_a_uniform_permutation() {
let mut rng = Rng::new(3);
let mut first = [0u32; 4];
for _ in 0..40_000 {
let mut v = [0, 1, 2, 3];
rng.shuffle(&mut v);
let mut sorted = v;
sorted.sort_unstable();
assert_eq!(sorted, [0, 1, 2, 3]);
first[v[0]] += 1;
}
assert!(first.iter().all(|&c| c.abs_diff(10_000) < 500), "{first:?}");
}
#[test]
fn keyed_normals_are_standard_normal() {
let key = stream_key(&[11, 0x5617_B000]);
let n = 200_000u64;
let draws: Vec<f64> = (0..n).map(|i| keyed_normal(key, i)).collect();
let mean = draws.iter().sum::<f64>() / n as f64;
let var = draws.iter().map(|z| (z - mean).powi(2)).sum::<f64>() / n as f64;
let tail = draws.iter().filter(|z| z.abs() > 1.959_964).count() as f64 / n as f64;
assert!(mean.abs() < 0.012, "mean {mean}");
assert!((var - 1.0).abs() < 0.016, "variance {var}");
assert!((tail - 0.05).abs() < 0.003, "two-sided 5% tail {tail}");
assert!(draws.iter().all(|z| z.is_finite()));
assert_ne!(
keyed_normal(key, 0),
keyed_normal(stream_key(&[12, 0x5617_B000]), 0)
);
}
}