use rand::RngCore;
use wide::f32x8;
use wide::f64x4;
use wide::i32x8;
use wide::u32x8;
use wide::u64x4;
use super::next_global_seed;
use super::xoshiro::F32_MAGIC;
use super::xoshiro::F64_MAGIC;
use super::xoshiro::Xoshiro128PP8;
use super::xoshiro::Xoshiro256PP4;
use super::xoshiro::splitmix64_next;
pub struct SimdRng {
pub(super) f64_engine: Xoshiro256PP4,
pub(super) f32_engine: Xoshiro128PP8,
u64_buf: [u64; 4],
u64_idx: usize,
f64_scalar_buf: [f64; 8],
f64_scalar_idx: usize,
f32_scalar_buf: [f32; 8],
f32_scalar_idx: usize,
i32_scalar_buf: [i32; 8],
i32_scalar_idx: usize,
}
impl SimdRng {
#[inline]
pub fn from_seed(seed: u64) -> Self {
let mut state = seed;
let seed64 = splitmix64_next(&mut state);
let seed32 = splitmix64_next(&mut state);
Self {
f64_engine: Xoshiro256PP4::new_from_u64(seed64),
f32_engine: Xoshiro128PP8::new_from_u64(seed32),
u64_buf: [0; 4],
u64_idx: 4,
f64_scalar_buf: [0.0; 8],
f64_scalar_idx: 8,
f32_scalar_buf: [0.0; 8],
f32_scalar_idx: 8,
i32_scalar_buf: [0; 8],
i32_scalar_idx: 8,
}
}
#[inline]
pub fn new() -> Self {
Self::from_seed(next_global_seed())
}
#[inline(always)]
pub fn next_i32x8(&mut self) -> i32x8 {
let raw = self.f32_engine.next();
unsafe { core::mem::transmute::<u32x8, i32x8>(raw) }
}
#[inline(always)]
pub fn next_f64_array(&mut self) -> [f64; 8] {
let a = self.f64_engine.next();
let b = self.f64_engine.next();
let magic = u64x4::splat(F64_MAGIC);
let one = f64x4::splat(1.0);
let bits_a = (a >> 12u32) | magic;
let bits_b = (b >> 12u32) | magic;
let fa: f64x4 = unsafe { core::mem::transmute::<u64x4, f64x4>(bits_a) };
let fb: f64x4 = unsafe { core::mem::transmute::<u64x4, f64x4>(bits_b) };
let ra = (fa - one).to_array();
let rb = (fb - one).to_array();
[ra[0], ra[1], ra[2], ra[3], rb[0], rb[1], rb[2], rb[3]]
}
#[inline(always)]
pub fn next_f32_array(&mut self) -> [f32; 8] {
let a = self.f32_engine.next();
let bits = (a >> 9u32) | u32x8::splat(F32_MAGIC);
let f: f32x8 = unsafe { core::mem::transmute::<u32x8, f32x8>(bits) };
(f - f32x8::splat(1.0)).to_array()
}
#[inline(always)]
pub fn next_f64(&mut self) -> f64 {
if self.f64_scalar_idx >= 8 {
let magic = u64x4::splat(F64_MAGIC);
let one = f64x4::splat(1.0);
let buf_ptr = self.f64_scalar_buf.as_mut_ptr();
unsafe {
let bits0 = (self.f64_engine.next() >> 12u32) | magic;
let f0: f64x4 = core::mem::transmute::<u64x4, f64x4>(bits0);
core::ptr::write_unaligned(buf_ptr as *mut f64x4, f0 - one);
let bits1 = (self.f64_engine.next() >> 12u32) | magic;
let f1: f64x4 = core::mem::transmute::<u64x4, f64x4>(bits1);
core::ptr::write_unaligned(buf_ptr.add(4) as *mut f64x4, f1 - one);
}
self.f64_scalar_idx = 0;
}
let v = self.f64_scalar_buf[self.f64_scalar_idx];
self.f64_scalar_idx += 1;
v
}
#[inline(always)]
pub fn next_f32(&mut self) -> f32 {
if self.f32_scalar_idx >= 8 {
let buf_ptr = self.f32_scalar_buf.as_mut_ptr();
unsafe {
let bits = (self.f32_engine.next() >> 9u32) | u32x8::splat(F32_MAGIC);
let f: f32x8 = core::mem::transmute::<u32x8, f32x8>(bits);
core::ptr::write_unaligned(buf_ptr as *mut f32x8, f - f32x8::splat(1.0));
}
self.f32_scalar_idx = 0;
}
let v = self.f32_scalar_buf[self.f32_scalar_idx];
self.f32_scalar_idx += 1;
v
}
#[inline(always)]
pub fn next_i32(&mut self) -> i32 {
if self.i32_scalar_idx >= 8 {
self.i32_scalar_buf = self.next_i32x8().to_array();
self.i32_scalar_idx = 0;
}
let v = self.i32_scalar_buf[self.i32_scalar_idx];
self.i32_scalar_idx += 1;
v
}
}
impl Default for SimdRng {
fn default() -> Self {
Self::new()
}
}
impl RngCore for SimdRng {
#[inline(always)]
fn next_u32(&mut self) -> u32 {
self.next_u64() as u32
}
#[inline(always)]
fn next_u64(&mut self) -> u64 {
let idx = self.u64_idx;
if idx >= 4 {
self.u64_buf = self.f64_engine.next().to_array();
self.u64_idx = 1;
return self.u64_buf[0];
}
self.u64_idx = idx + 1;
self.u64_buf[idx]
}
fn fill_bytes(&mut self, dest: &mut [u8]) {
let mut written = 0;
let total = dest.len();
while self.u64_idx < 4 && total - written >= 8 {
let v = self.u64_buf[self.u64_idx];
self.u64_idx += 1;
dest[written..written + 8].copy_from_slice(&v.to_le_bytes());
written += 8;
}
while total - written >= 32 {
let block = self.f64_engine.next().to_array();
dest[written..written + 8].copy_from_slice(&block[0].to_le_bytes());
dest[written + 8..written + 16].copy_from_slice(&block[1].to_le_bytes());
dest[written + 16..written + 24].copy_from_slice(&block[2].to_le_bytes());
dest[written + 24..written + 32].copy_from_slice(&block[3].to_le_bytes());
written += 32;
}
if written == total {
return;
}
self.u64_buf = self.f64_engine.next().to_array();
self.u64_idx = 0;
while total - written >= 8 {
let v = self.u64_buf[self.u64_idx];
self.u64_idx += 1;
dest[written..written + 8].copy_from_slice(&v.to_le_bytes());
written += 8;
}
if written < total {
let bytes = self.u64_buf[self.u64_idx].to_le_bytes();
let take = total - written;
dest[written..written + take].copy_from_slice(&bytes[..take]);
self.u64_idx += 1;
}
}
}