#![allow(unsafe_code)]
use crate::rng::SplitMix64;
const ZIG_N: usize = 128;
const ZIG_R: f64 = 3.442_619_855_899;
include!("ziggurat_tables.rs");
fn ziggurat_scalar(rng: &mut SplitMix64) -> f64 {
loop {
let word = rng.next_u64();
let (sign, layer, j) = unpack(word);
let kj = ZIG_K.get(layer).copied().unwrap_or(0);
let wj = ZIG_W.get(layer).copied().unwrap_or(0.0);
if j < kj {
return sign * f64::from(j) * wj;
}
if let Some(v) = ziggurat_fixup(rng, sign, layer, j) {
return v;
}
}
}
fn unpack(word: u64) -> (f64, usize, u32) {
let low = word & 0xFFFF_FFFF;
let sign = if word & 0x1_0000_0000 == 0 { 1.0 } else { -1.0 };
let layer = usize::try_from(low).unwrap_or(0) & (ZIG_N - 1);
let j = u32::try_from(low & 0x7FFF_FFFF).unwrap_or(0);
(sign, layer, j)
}
fn ziggurat_fixup(rng: &mut SplitMix64, sign: f64, layer: usize, j: u32) -> Option<f64> {
if layer == 0 {
loop {
let x = -rng.next_f64().max(f64::MIN_POSITIVE).ln() / ZIG_R;
let y = -rng.next_f64().max(f64::MIN_POSITIVE).ln();
if y + y > x * x {
return Some(sign * (ZIG_R + x));
}
}
}
let wj = ZIG_W.get(layer).copied().unwrap_or(0.0);
let x = f64::from(j) * wj;
let f_lo = ZIG_F.get(layer).copied().unwrap_or(0.0);
let f_hi = ZIG_F.get(layer - 1).copied().unwrap_or(0.0);
let u = rng.next_f64();
if u.mul_add(f_hi - f_lo, f_lo) < (-0.5 * x * x).exp() {
Some(sign * x)
} else {
None
}
}
pub(super) fn normal_sample_into(mean: f64, std_dev: f64, rng: &mut SplitMix64, out: &mut [f64]) {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe {
sample_neon(mean, std_dev, rng, out);
}
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
unsafe {
sample_avx2(mean, std_dev, rng, out);
}
return;
}
}
sample_scalar(mean, std_dev, rng, out);
}
fn sample_scalar(mean: f64, std_dev: f64, rng: &mut SplitMix64, out: &mut [f64]) {
for slot in out.iter_mut() {
*slot = std_dev.mul_add(ziggurat_scalar(rng), mean);
}
}
fn fast_candidate(word: u64) -> (f64, bool) {
let (sign, layer, j) = unpack(word);
let kj = ZIG_K.get(layer).copied().unwrap_or(0);
let wj = ZIG_W.get(layer).copied().unwrap_or(0.0);
(sign * f64::from(j) * wj, j < kj)
}
fn resolve_lane(word: u64, rng: &mut SplitMix64) -> f64 {
let (candidate, accepted) = fast_candidate(word);
if accepted {
candidate
} else {
let (sign, layer, j) = unpack(word);
ziggurat_fixup_or_retry(rng, sign, layer, j)
}
}
fn ziggurat_fixup_or_retry(rng: &mut SplitMix64, sign: f64, layer: usize, j: u32) -> f64 {
if let Some(v) = ziggurat_fixup(rng, sign, layer, j) {
return v;
}
ziggurat_scalar(rng)
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn sample_neon(mean: f64, std_dev: f64, rng: &mut SplitMix64, out: &mut [f64]) {
use std::arch::aarch64::{vfmaq_f64, vld1q_f64, vsetq_lane_f64, vst1q_f64};
let lanes = 2;
let n = out.len();
let body = n - (n % lanes);
let vmean = std::arch::aarch64::vdupq_n_f64(mean);
let vstd = std::arch::aarch64::vdupq_n_f64(std_dev);
let mut i = 0;
while i < body {
let w0 = rng.next_u64();
let w1 = rng.next_u64();
let (c0, a0) = fast_candidate(w0);
let (c1, a1) = fast_candidate(w1);
let z0 = if a0 { c0 } else { resolve_lane(w0, rng) };
let z1 = if a1 { c1 } else { resolve_lane(w1, rng) };
let mut zv = vsetq_lane_f64::<0>(z0, vld1q_f64([0.0_f64, 0.0].as_ptr()));
zv = vsetq_lane_f64::<1>(z1, zv);
let res = vfmaq_f64(vmean, vstd, zv);
vst1q_f64(out.as_mut_ptr().add(i), res);
i += lanes;
}
if body < n {
sample_scalar(mean, std_dev, rng, out.get_mut(body..n).unwrap_or(&mut []));
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn sample_avx2(mean: f64, std_dev: f64, rng: &mut SplitMix64, out: &mut [f64]) {
use std::arch::x86_64::{_mm256_fmadd_pd, _mm256_loadu_pd, _mm256_set1_pd, _mm256_storeu_pd};
let lanes = 4;
let n = out.len();
let body = n - (n % lanes);
let vmean = _mm256_set1_pd(mean);
let vstd = _mm256_set1_pd(std_dev);
let mut i = 0;
while i < body {
let words = [
rng.next_u64(),
rng.next_u64(),
rng.next_u64(),
rng.next_u64(),
];
let mut zs = [0.0_f64; 4];
for (slot, &word) in zs.iter_mut().zip(words.iter()) {
let (cand, ok) = fast_candidate(word);
*slot = if ok { cand } else { resolve_lane(word, rng) };
}
let zv = unsafe { _mm256_loadu_pd(zs.as_ptr()) };
let res = _mm256_fmadd_pd(vstd, zv, vmean);
unsafe {
_mm256_storeu_pd(out.as_mut_ptr().add(i), res);
}
i += lanes;
}
if body < n {
sample_scalar(mean, std_dev, rng, out.get_mut(body..n).unwrap_or(&mut []));
}
}
#[cfg(test)]
mod tests {
use super::*;
const TWO_POW_31: f64 = 2_147_483_648.0;
#[test]
fn table_construction_invariant_holds() {
let top = ZIG_F.first().copied().unwrap_or(0.0);
assert!(
(top - 1.0).abs() < 1e-15,
"top-layer density must be 1, was {top}"
);
for (i, (&w, &fi)) in ZIG_W.iter().zip(ZIG_F.iter()).enumerate().skip(1) {
let edge = w * TWO_POW_31;
let want = (-0.5 * edge * edge).exp();
assert!(
(fi - want).abs() < 1e-12,
"layer {i}: ZIG_F={fi} != f(edge)={want}"
);
}
let bottom_edge = ZIG_W.last().copied().unwrap_or(0.0) * TWO_POW_31;
assert!(
(bottom_edge - ZIG_R).abs() < 1e-9,
"bottom edge {bottom_edge} must equal R={ZIG_R}"
);
}
fn normal_cdf(x: f64) -> f64 {
0.5 * (1.0 + crate::special::erf(x / std::f64::consts::SQRT_2))
}
fn ks_statistic(sorted: &[f64], cdf: impl Fn(f64) -> f64) -> f64 {
let n = f64::from(u32::try_from(sorted.len()).unwrap_or(u32::MAX));
let mut d = 0.0_f64;
for (i, &x) in sorted.iter().enumerate() {
let f = cdf(x);
let i_f = f64::from(u32::try_from(i).unwrap_or(u32::MAX));
d = d.max((i_f + 1.0) / n - f).max(f - i_f / n);
}
d
}
#[test]
fn scalar_ziggurat_fits_standard_normal() {
let mut rng = SplitMix64::new(20_240_628);
let n = 100_000usize;
let mut xs: Vec<f64> = (0..n).map(|_| ziggurat_scalar(&mut rng)).collect();
let count = f64::from(u32::try_from(n).unwrap_or(u32::MAX));
let mean = xs.iter().sum::<f64>() / count;
let var = xs.iter().map(|x| (x - mean) * (x - mean)).sum::<f64>() / count;
assert!(mean.abs() < 0.02, "empirical mean {mean} not near 0");
assert!(
(var - 1.0).abs() < 0.03,
"empirical variance {var} not near 1"
);
xs.sort_by(f64::total_cmp);
let ks = ks_statistic(&xs, normal_cdf);
let crit = 1.63 / count.sqrt();
assert!(ks < crit, "KS={ks} exceeds 1% critical {crit}");
}
#[test]
fn batch_fill_is_reproducible_and_affine() {
let (mean, std_dev) = (1.5, 2.0);
let fill = |seed: u64, len: usize| {
let mut rng = SplitMix64::new(seed);
let mut out = vec![0.0; len];
normal_sample_into(mean, std_dev, &mut rng, &mut out);
out
};
let a = fill(424_242, 103);
assert_eq!(a, fill(424_242, 103), "batch fill not reproducible");
let big = fill(99, 80_000);
let count = f64::from(u32::try_from(big.len()).unwrap_or(u32::MAX));
let emp_mean = big.iter().sum::<f64>() / count;
let emp_var = big
.iter()
.map(|x| (x - emp_mean) * (x - emp_mean))
.sum::<f64>()
/ count;
assert!(
(emp_mean - mean).abs() < 0.05,
"mean {emp_mean} not near {mean}"
);
let want_var = std_dev * std_dev;
assert!(
(emp_var - want_var).abs() < 0.15,
"variance {emp_var} not near {want_var}"
);
}
}