use std::f32::consts::SQRT_2;
const NEXT_BELOW_ONE: f32 = 0.99999994_f32;
const NEXT_ABOVE_NEG_ONE: f32 = -0.99999994_f32;
const U32_MAX_AS_F32: f32 = 4294967295.0_f32;
#[must_use]
pub fn key(seed: u64) -> [u32; 2] {
let k1 = (seed >> 32) as u32;
let k2 = seed as u32;
[k1, k2]
}
#[must_use]
pub fn threefry2x32(key: [u32; 2], count: [u32; 2]) -> [u32; 2] {
const ROTATIONS: [[u32; 4]; 2] = [[13, 15, 26, 6], [17, 29, 16, 24]];
const PARITY: u32 = 0x1BD1_1BDA;
let ks: [u32; 3] = [key[0], key[1], key[0] ^ key[1] ^ PARITY];
let mut c0 = count[0].wrapping_add(ks[0]);
let mut c1 = count[1].wrapping_add(ks[1]);
for i in 0..5usize {
for &r in &ROTATIONS[i % 2] {
c0 = c0.wrapping_add(c1);
c1 = c1.rotate_left(r) ^ c0;
}
c0 = c0.wrapping_add(ks[(i + 1) % 3]);
c1 = c1
.wrapping_add(ks[(i + 2) % 3])
.wrapping_add((i as u32).wrapping_add(1));
}
[c0, c1]
}
#[must_use]
pub fn random_bits(n: usize, key: [u32; 2]) -> Vec<u32> {
let mut out = vec![0u32; n];
if n == 0 {
return out;
}
let out_skip = n;
let half = out_skip / 2;
let even = out_skip % 2 == 0;
let mut c_first: u32 = 0;
let mut c_second: u32 = half as u32 + u32::from(!even);
while (c_first as usize) + 1 < half {
let rb = threefry2x32(key, [c_first, c_second]);
out[c_first as usize] = rb[0];
out[c_second as usize] = rb[1];
c_first = c_first.wrapping_add(1);
c_second = c_second.wrapping_add(1);
}
if (c_first as usize) < half {
let rb = threefry2x32(key, [c_first, c_second]);
out[c_first as usize] = rb[0];
c_first = c_first.wrapping_add(1);
if (c_second as usize) < n {
out[c_second as usize] = rb[1];
}
}
if !even {
let rb = threefry2x32(key, [half as u32, 0]);
if half < n {
out[half] = rb[0];
}
}
let _ = c_first;
out
}
#[must_use]
pub fn uniform(n: usize, key: [u32; 2], lo: f32, hi: f32) -> Vec<f32> {
let bits = random_bits(n, key);
let range = hi - lo;
bits.into_iter()
.map(|b| {
let u = (b as f32) / U32_MAX_AS_F32;
let u = if u < NEXT_BELOW_ONE {
u
} else {
NEXT_BELOW_ONE
};
range * u + lo
})
.collect()
}
#[must_use]
pub fn normal(n: usize, key: [u32; 2]) -> Vec<f32> {
let u = uniform(n, key, NEXT_ABOVE_NEG_ONE, 1.0_f32);
u.into_iter().map(|x| SQRT_2 * erfinv(x)).collect()
}
#[allow(clippy::excessive_precision)]
#[must_use]
pub fn erfinv(a: f32) -> f32 {
let t = a.mul_add(-a, 1.0_f32).ln();
let lhs = |t: f32| -> f32 {
let mut p = 3.03697567e-10_f32;
p = p.mul_add(t, 2.93243101e-8_f32);
p = p.mul_add(t, 1.22150334e-6_f32);
p = p.mul_add(t, 2.84108955e-5_f32);
p = p.mul_add(t, 3.93552968e-4_f32);
p = p.mul_add(t, 3.02698812e-3_f32);
p = p.mul_add(t, 4.83185798e-3_f32);
p = p.mul_add(t, -2.64646143e-1_f32);
p.mul_add(t, 8.40016484e-1_f32)
};
let rhs = |t: f32| -> f32 {
let mut p = 5.43877832e-9_f32;
p = p.mul_add(t, 1.43285448e-7_f32);
p = p.mul_add(t, 1.22774793e-6_f32);
p = p.mul_add(t, 1.12963626e-7_f32);
p = p.mul_add(t, -5.61530760e-5_f32);
p = p.mul_add(t, -1.47697632e-4_f32);
p = p.mul_add(t, 2.31468678e-3_f32);
p = p.mul_add(t, 1.15392581e-2_f32);
p = p.mul_add(t, -2.32015476e-1_f32);
p.mul_add(t, 8.86226892e-1_f32)
};
let thresh = 6.125_f32;
let p = if t.abs() > thresh { lhs(t) } else { rhs(t) };
a * p
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn key_vectors() {
assert_eq!(key(0), [0, 0]);
assert_eq!(key(1), [0, 1]);
let seed = 1u64 << 32;
assert_eq!(key(seed), [1, 0]);
assert_eq!(key(seed + 1), [1, 1]);
}
#[test]
fn split_two_matches_mlx() {
let v = random_bits(4, key(0));
assert_eq!(v, vec![4146024105, 967050713, 2718843009, 1272950319]);
}
#[test]
fn split_three_matches_mlx() {
let v = random_bits(6, key(0));
assert_eq!(
v,
vec![2467461003, 428148500, 3186719485, 3840466878, 2562233961, 1946702221]
);
}
#[test]
fn scalar_bits_matches_mlx() {
assert_eq!(random_bits(1, key(0)), vec![1797259609]);
assert_eq!(random_bits(1, key(1)), vec![507451445]);
}
#[test]
fn three_bits_odd_layout_matches_mlx() {
assert_eq!(
random_bits(3, key(0)),
vec![4146024105, 1351547692, 2718843009]
);
}
#[test]
fn threefry_pair_equals_bits_two() {
let rb = threefry2x32(key(0), [0, 1]);
assert_eq!(random_bits(2, key(0)), vec![rb[0], rb[1]]);
}
#[test]
fn bits_four_first_word_uses_count_zero_two() {
let rb = threefry2x32(key(0), [0, 2]);
let v = random_bits(4, key(0));
assert_eq!(v[0], rb[0]); assert_eq!(v[2], rb[1]); assert_eq!(v[0], 4146024105); }
#[test]
fn uniform_scalar_matches_mlx() {
let u0 = uniform(1, key(0), 0.0, 1.0)[0];
let expected0 = 1797259609.0_f32 / U32_MAX_AS_F32;
assert_eq!(u0, expected0);
let u1 = uniform(1, key(1), 0.0, 1.0)[0];
let expected1 = 507451445.0_f32 / U32_MAX_AS_F32;
assert_eq!(u1, expected1);
}
#[test]
fn uniform_upper_bound_respected() {
let v = uniform(4096, key(128291), -1.0, 1.0);
assert!(v.iter().all(|&x| x < 1.0));
assert!(v.iter().all(|&x| x >= -1.0));
}
#[test]
fn erfinv_basic() {
assert_eq!(erfinv(0.0), 0.0);
let x = 0.3_f32;
assert!((erfinv(-x) + erfinv(x)).abs() < 1e-6);
}
#[test]
fn normal_finite() {
let v = normal(131072, key(42));
assert!(v.iter().all(|x| x.is_finite()));
}
}