use wrapn::{wrap, wu64};
use crate::{_internal::sm64_from_seed32, Rng};
#[repr(C, align(64))]
#[derive(Debug, Clone, Copy)]
pub struct Pcg32 {
state: wu64,
inc: wu64,
}
impl Pcg32 {
pub const fn new(seed: u32) -> Self {
let mut seedgen = sm64_from_seed32!(seed);
Pcg32 {
state: wrap!(seedgen.nextu_const()),
inc: wrap!(seedgen.nextu_const()),
}
}
}
impl Rng for Pcg32 {
type Word = u32;
#[inline]
fn nextu(&mut self) -> Self::Word {
let oldstate = self.state;
self.state = oldstate * 6364136223846793005 + self.inc;
let xorshifted = (((oldstate >> 18) ^ oldstate) >> 27).cast::<u32>();
let rot = (oldstate >> 59).cast::<u32>();
*xorshifted.rotate_right(*rot)
}
}
#[cfg(feature = "simd")]
pub use simd::*;
#[cfg(feature = "simd")]
pub mod simd {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
use crate::RngV;
pub const PCG32X8_LANE: usize = 8;
pub const PCG32X8_PAR_CHUNK: usize = 131_072;
pub const PCG32X8_PAR_CHUNK_BLOCKS: u64 = (PCG32X8_PAR_CHUNK / PCG32X8_LANE) as u64;
pub const PCG32_MULT: u64 = 6364136223846793005;
#[cfg(target_arch = "x86_64")]
#[repr(C, align(64))]
pub struct Pcg32x8 {
pub(crate) state: __m512i,
pub(crate) inc: __m512i,
}
#[cfg(target_arch = "x86_64")]
impl Pcg32x8 {
#[target_feature(enable = "avx512f")]
pub unsafe fn new(seed: u32) -> Self {
let mut seedgen = crate::sm64_from_seed32!(seed);
let mut state = [0u64; PCG32X8_LANE];
state.iter_mut().for_each(|v| *v = seedgen.nextu_const());
let mut inc = [0u64; PCG32X8_LANE];
inc.iter_mut().for_each(|v| *v = seedgen.nextu_const());
unsafe {
Pcg32x8 {
state: _mm512_loadu_si512(state.as_ptr() as _),
inc: _mm512_loadu_si512(inc.as_ptr() as _),
}
}
}
#[allow(unsafe_op_in_unsafe_fn)]
#[target_feature(enable = "avx512f")]
pub unsafe fn nextu(&mut self) -> [u32; PCG32X8_LANE] {
let mult_lo = _mm512_set1_epi64(0x4C957F2D_i64);
let mult_hi = _mm512_set1_epi64(0x5851F42D_i64);
let mask32 = _mm512_set1_epi64(0xFFFFFFFF_i64);
let out256 = Self::step_u32(&mut self.state, self.inc, mult_lo, mult_hi, mask32);
std::mem::transmute(out256)
}
#[cfg(target_arch = "x86_64")]
#[inline]
#[allow(unsafe_op_in_unsafe_fn)]
#[target_feature(enable = "avx512f")]
pub(crate) unsafe fn step_u32(
state: &mut __m512i,
inc: __m512i,
mult_lo: __m512i,
mult_hi: __m512i,
mask32: __m512i,
) -> __m256i {
let oldstate = *state;
let state_hi = _mm512_srli_epi64(oldstate, 32);
let prod_lo = _mm512_mul_epu32(oldstate, mult_lo);
let cross = _mm512_add_epi64(
_mm512_mul_epu32(state_hi, mult_lo),
_mm512_mul_epu32(oldstate, mult_hi),
);
*state = _mm512_add_epi64(_mm512_add_epi64(prod_lo, _mm512_slli_epi64(cross, 32)), inc);
let xs = _mm512_srli_epi64(
_mm512_xor_si512(_mm512_srli_epi64(oldstate, 18), oldstate),
27,
);
let rot = _mm512_srli_epi64(oldstate, 59);
let rotated = _mm512_rorv_epi32(_mm512_and_si512(xs, mask32), rot);
_mm512_cvtepi64_epi32(rotated)
}
}
impl RngV for Pcg32x8 {
type Word = __m256i;
#[inline]
fn nextuv(&mut self) -> Self::Word {
unsafe {
let mult_lo = _mm512_set1_epi64(0x4C957F2D_i64);
let mult_hi = _mm512_set1_epi64(0x5851F42D_i64);
let mask32 = _mm512_set1_epi64(0xFFFFFFFF_i64);
Self::step_u32(&mut self.state, self.inc, mult_lo, mult_hi, mask32)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
crate::safe_test! { Pcg32 }
}