urng 1.0.0

Universal Random Number Generator
use wrapn::{wrap, wu64};

use crate::prng::b64::SplitMix64;
use crate::rng::Rng;

/// CET 64-bit RNG implementation.
///
/// # Example
/// ```
/// use urng::{Rng, Cet64};
///
/// let mut rng = Cet64::new(1);
/// let _ = rng.nextu();
/// ```
#[repr(C, align(64))]
#[derive(Debug, Clone, Copy)]
pub struct Cet64 {
    s: wu64,
}

const SP1: u64 = 0xFFFFFFFFFFFFFF43;
const SP2: u64 = 0xFFFFFFFFFFFFFF1B;
const P1: u64 = 0x94D049BB133111EB;

impl Cet64 {
    /// Creates a new `Cet64` instance with the given seed.
    pub const fn new(seed: u64) -> Self {
        let mut seedgen = SplitMix64::new(seed);
        Self {
            s: wrap!(seedgen.nextu_const()),
        }
    }
}

impl Rng for Cet64 {
    type Word = u64;

    #[inline(always)]
    fn nextu(&mut self) -> Self::Word {
        self.s += SP1;

        let mut x = self.s;
        x ^= x >> 30;
        x *= SP2;
        x ^= x >> 27;
        x *= P1;
        x ^= x >> 31;

        *x
    }
}

/// CET 256-bit-state 64-bit RNG implementation.
///
/// # Example
/// ```
/// use urng::{Rng, Cet256};
///
/// let mut rng = Cet256::new(1);
/// let _ = rng.nextu();
/// ```
#[repr(C, align(64))]
#[derive(Debug, Clone, Copy)]
pub struct Cet256 {
    s: [wu64; 4],
}

impl Cet256 {
    /// Creates a new `Cet256` instance with the given seed.
    pub const fn new(seed: u64) -> Self {
        let mut seedgen = SplitMix64::new(seed);
        Self {
            s: wrap![
                seedgen.nextu_const(),
                seedgen.nextu_const(),
                seedgen.nextu_const(),
                seedgen.nextu_const(),
            ],
        }
    }
}

impl Rng for Cet256 {
    type Word = u64;

    #[inline(always)]
    fn nextu(&mut self) -> Self::Word {
        self.s[0] += SP1;
        let c0 = (self.s[0] < SP1) as u64;
        self.s[1] += c0;
        let c1 = (self.s[1] < c0) as u64;
        self.s[2] += c1;
        let c2 = (self.s[2] < c1) as u64;
        self.s[3] += c2;

        let mut x = self.s[0] ^ self.s[3];
        x += self.s[1].rotate_left(17);

        x ^= x >> 30;
        x *= SP2;
        x ^= x >> 27;
        x *= P1;
        x ^= x >> 31;

        *x
    }
}

#[cfg(feature = "simd")]
pub use simd::*;

#[cfg(feature = "simd")]
pub mod simd {
    use std::arch::x86_64::*;

    use crate::_internal::{i2f_bits, u2f_01};
    use crate::SplitMix64;

    use crate::prng::b64::cet::{P1, SP1};

    /// An 8-way SIMD CET64 generator using AVX-512 512-bit intrinsics.
    #[cfg(target_arch = "x86_64")]
    #[repr(C, align(64))]
    pub struct Cet64x8 {
        s: __m512i,
    }

    #[cfg(target_arch = "x86_64")]
    impl Cet64x8 {
        /// Creates a new `Cet64x8` from 8 independent seeds.
        pub fn new(seed: u64) -> Self {
            let mut seedgen = SplitMix64::new(seed);
            let s = [0u64; 8].map(|_| seedgen.nextu_const());
            unsafe {
                Self {
                    s: _mm512_loadu_si512(s.as_ptr() as *const _),
                }
            }
        }

        #[target_feature(enable = "avx512f")]
        #[allow(unsafe_op_in_unsafe_fn)]
        pub(crate) unsafe fn mul_sp2_vec(x: __m512i) -> __m512i {
            let x8 = _mm512_slli_epi64(x, 8);
            let x5 = _mm512_slli_epi64(x, 5);
            let x2 = _mm512_slli_epi64(x, 2);
            let t0 = _mm512_sub_epi64(x8, x5);
            let t1 = _mm512_add_epi64(x2, x);
            let t = _mm512_add_epi64(t0, t1);
            _mm512_sub_epi64(_mm512_setzero_si512(), t)
        }

        /// Generates 8 random `u64` values simultaneously.
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        #[allow(unsafe_op_in_unsafe_fn)]
        pub unsafe fn nextu(&mut self) -> [u64; 8] {
            let sp1 = _mm512_set1_epi64(SP1 as i64);
            let p1 = _mm512_set1_epi64(P1 as i64);

            self.s = _mm512_add_epi64(self.s, sp1);

            let mut x = self.s;
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 30));
            x = Self::mul_sp2_vec(x);
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 27));
            x = _mm512_mullo_epi64(x, p1);
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 31));

            let mut out = [0u64; 8];
            _mm512_storeu_si512(out.as_mut_ptr() as *mut __m512i, x);
            out
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn nextf(&mut self) -> [f64; 8] {
            let u = unsafe { self.nextu() };
            let mut out = [0f64; 8];
            for i in 0..8 {
                out[i] = u2f_01!(f64, 64, u[i]);
            }
            out
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn randi(&mut self, min: i64, max: i64) -> [i64; 8] {
            let u = unsafe { self.nextu() };
            let range = (max as i128 - min as i128 + 1) as u128;
            let mut out = [0i64; 8];
            for i in 0..8 {
                out[i] = ((u[i] as u128 * range) >> 64) as i64 + min;
            }
            out
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn randf(&mut self, min: f64, max: f64) -> [f64; 8] {
            let u = unsafe { self.nextu() };
            let range = max - min;
            let mut out = [0f64; 8];
            for i in 0..8 {
                out[i] = u2f_01!(f64, 64, u[i]) * range + min;
            }
            out
        }
    }

    /// A 2-way SIMD CET256 generator using AVX-512 storage layout.
    ///
    /// Internal lane mapping is `[s0, s1, s2, s3, s0, s1, s2, s3]`.
    #[cfg(target_arch = "x86_64")]
    #[repr(C, align(64))]
    pub struct Cet256x2 {
        s: __m512i,
    }

    #[cfg(target_arch = "x86_64")]
    impl Cet256x2 {
        /// Creates a new `Cet256x2` from 2 independent seeds.
        pub fn new(seed: u64) -> Self {
            let mut seedgen = SplitMix64::new(seed);
            let mut s = [0u64; 8];
            for lane in &mut s {
                *lane = seedgen.nextu_const();
            }
            unsafe {
                Self {
                    s: _mm512_loadu_si512(s.as_ptr() as *const _),
                }
            }
        }

        /// Generates 2 random `u64` values simultaneously.
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        #[allow(unsafe_op_in_unsafe_fn)]
        pub unsafe fn nextu(&mut self) -> [u64; 2] {
            let mut state = [0u64; 8];
            _mm512_storeu_si512(state.as_mut_ptr() as *mut __m512i, self.s);

            let mut lanes = [0u64; 8];
            for base in [0usize, 4usize] {
                state[base] = state[base].wrapping_add(SP1);
                let c0 = (state[base] < SP1) as u64;
                state[base + 1] = state[base + 1].wrapping_add(c0);
                let c1 = (state[base + 1] < c0) as u64;
                state[base + 2] = state[base + 2].wrapping_add(c1);
                let c2 = (state[base + 2] < c1) as u64;
                state[base + 3] = state[base + 3].wrapping_add(c2);

                lanes[base] =
                    (state[base] ^ state[base + 3]).wrapping_add(state[base + 1].rotate_left(17));
            }

            self.s = _mm512_loadu_si512(state.as_ptr() as *const __m512i);

            let p1 = _mm512_set1_epi64(P1 as i64);
            let mut x = _mm512_loadu_si512(lanes.as_ptr() as *const __m512i);
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 30));
            x = Cet64x8::mul_sp2_vec(x);
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 27));
            x = _mm512_mullo_epi64(x, p1);
            x = _mm512_xor_si512(x, _mm512_srli_epi64(x, 31));

            let mut out_lanes = [0u64; 8];
            _mm512_storeu_si512(out_lanes.as_mut_ptr() as *mut __m512i, x);
            [out_lanes[0], out_lanes[4]]
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn nextf(&mut self) -> [f64; 2] {
            let u = unsafe { self.nextu() };
            [u2f_01!(f64, 64, u[0]), u2f_01!(f64, 64, u[1])]
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn randi(&mut self, min: i64, max: i64) -> [i64; 2] {
            let u = unsafe { self.nextu() };
            let range = (max as i128 - min as i128 + 1) as u128;
            [
                ((u[0] as u128 * range) >> 64) as i64 + min,
                ((u[1] as u128 * range) >> 64) as i64 + min,
            ]
        }

        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn randf(&mut self, min: f64, max: f64) -> [f64; 2] {
            let u = unsafe { self.nextu() };
            let range = max - min;
            [
                u2f_01!(f64, 64, u[0]) * range + min,
                u2f_01!(f64, 64, u[1]) * range + min,
            ]
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    crate::safe_test! {
        Cet64,
        Cet256
    }

    #[cfg(all(
        feature = "simd",
        target_feature = "avx512f",
        target_feature = "avx512dq"
    ))]
    crate::unsafe_test! {
        Cet64x8,
        Cet256x2
    }
}