use std::slice::from_raw_parts_mut;
use crate::prng::b32::SplitMix32;
use crate::rng::Rng;
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_new(seed: u32) -> *mut SplitMix32 {
Box::into_raw(Box::new(SplitMix32::new(seed)))
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_free(ptr: *mut SplitMix32) {
if !ptr.is_null() {
unsafe {
let _ = Box::from_raw(ptr);
}
}
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_next_u32s(ptr: *mut SplitMix32, out: *mut u32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
crate::_internal::par_fill_reseed32(buffer, rng.nextu(), SplitMix32::new, |r| r.nextu());
}
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_next_f32s(ptr: *mut SplitMix32, out: *mut f32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
crate::_internal::par_fill_reseed32(buffer, rng.nextu(), SplitMix32::new, |r| r.nextf());
}
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_rand_i32s(
ptr: *mut SplitMix32,
out: *mut i32,
count: usize,
min: i32,
max: i32,
) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
crate::_internal::par_fill_reseed32(buffer, rng.nextu(), SplitMix32::new, |r| {
r.randi(min, max)
});
}
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32_rand_f32s(
ptr: *mut SplitMix32,
out: *mut f32,
count: usize,
min: f32,
max: f32,
) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
crate::_internal::par_fill_reseed32(buffer, rng.nextu(), SplitMix32::new, |r| {
r.randf(min, max)
});
}
}
#[cfg(feature = "simd")]
pub use simd::*;
#[cfg(feature = "simd")]
mod simd {
use super::*;
use crate::prng::b32::{
SPLITMIX32_GAMMA, SPLITMIX32X16, SPLITMIX32X16_PAR_CHUNK, SplitMix32x16,
};
use rayon::iter::{IndexedParallelIterator, ParallelIterator};
use rayon::slice::ParallelSliceMut;
use std::arch::x86_64::*;
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32x16_new(seed: u32) -> *mut SplitMix32x16 {
unsafe { Box::into_raw(Box::new(SplitMix32x16::new(seed))) }
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32x16_free(ptr: *mut SplitMix32x16) {
if !ptr.is_null() {
unsafe {
let _ = Box::from_raw(ptr);
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
#[allow(unsafe_op_in_unsafe_fn)]
unsafe fn splitmix32x16_next_u32s_chunk(chunk_idx: usize, chunk: &mut [u32], state0: __m512i) {
let offset = ((chunk_idx * SPLITMIX32X16_PAR_CHUNK) as u32).wrapping_mul(SPLITMIX32_GAMMA);
let mut state = _mm512_add_epi32(state0, _mm512_set1_epi32(offset as i32));
let step = _mm512_set1_epi32(SPLITMIX32_GAMMA.wrapping_mul(SPLITMIX32X16 as u32) as i32);
let is_aligned = (chunk.as_ptr() as usize) & 63 == 0;
let mut chunks16 = chunk.chunks_exact_mut(SPLITMIX32X16);
if is_aligned {
for dst in chunks16.by_ref() {
let v = SplitMix32x16::compute(state);
_mm512_stream_si512(dst.as_mut_ptr() as *mut _, v);
state = _mm512_add_epi32(state, step);
}
} else {
for dst in chunks16.by_ref() {
let v = SplitMix32x16::compute(state);
_mm512_storeu_si512(dst.as_mut_ptr() as *mut _, v);
state = _mm512_add_epi32(state, step);
}
}
let rem = chunks16.into_remainder();
if !rem.is_empty() {
let v = SplitMix32x16::compute(state);
let mut tmp = [0u32; SPLITMIX32X16];
_mm512_storeu_si512(tmp.as_mut_ptr() as *mut _, v);
rem.copy_from_slice(&tmp[..rem.len()]);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn splitmix32x16_next_u32s(
ptr: *mut SplitMix32x16,
out: *mut u32,
count: usize,
) {
if count == 0 {
return;
}
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
let state0 = rng.state;
buffer
.par_chunks_mut(SPLITMIX32X16_PAR_CHUNK)
.enumerate()
.for_each(|(chunk_idx, chunk)| {
splitmix32x16_next_u32s_chunk(chunk_idx, chunk, state0);
});
rng.state = _mm512_add_epi32(
rng.state,
_mm512_set1_epi32((count as u32).wrapping_mul(SPLITMIX32_GAMMA) as i32),
);
}
}
}