use std::slice::from_raw_parts_mut;
use rayon::iter::{IndexedParallelIterator, ParallelIterator};
use rayon::slice::ParallelSliceMut;
use wrapn::wu32;
use crate::_internal::{fill_chunk_auto, prefer_nt};
use crate::cbrng::b32::{Threefry32x2, Threefry32x4};
use crate::{i2f_bits, u2f_01};
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_new(seed: u32) -> *mut Threefry32x4 {
Box::into_raw(Box::new(Threefry32x4::new(seed)))
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_free(ptr: *mut Threefry32x4) {
if !ptr.is_null() {
unsafe {
let _ = Box::from_raw(ptr);
}
}
}
const THREEFRY32_PAR_CHUNK: usize = 0x20000;
#[inline(always)]
fn fry4_fill<T, M>(buffer: &mut [T], c0: [wu32; 4], k: [wu32; 5], tw: [wu32; 3], map: M)
where
T: Copy + Default + Send,
M: Fn(wu32) -> T + Sync,
{
let nt = prefer_nt::<T>(buffer.len());
buffer
.par_chunks_mut(THREEFRY32_PAR_CHUNK)
.enumerate()
.for_each(|(chunk_idx, chunk)| {
let chunk_base = (chunk_idx * (THREEFRY32_PAR_CHUNK >> 2)) as u64;
let c0_64 = (c0[0].cast::<u64>()) | (c0[1].cast::<u64>() << 32);
let mut c64 = c0_64 + chunk_base;
unsafe {
fill_chunk_auto(chunk, nt, || {
let mut out = [T::default(); 64];
for b in 0..16 {
let cc = c64 + (b as u64);
let c = [cc.cast(), (cc >> 32).cast(), c0[2], c0[3]];
let r = Threefry32x4::compute(c, &k, &tw);
out[b * 4] = map(r[0]);
out[b * 4 + 1] = map(r[1]);
out[b * 4 + 2] = map(r[2]);
out[b * 4 + 3] = map(r[3]);
}
c64 += 16;
out
});
}
});
}
#[inline(always)]
fn fry4_advance(rng: &mut Threefry32x4, count: usize) {
let num_blocks = (count.div_ceil(64) * 16) as u64;
let c0_64 = (rng.c[0].cast::<u64>()) | ((rng.c[1].cast::<u64>()) << 32);
let new_c64 = c0_64 + num_blocks;
rng.c[0] = new_c64.cast();
rng.c[1] = (new_c64 >> 32).cast();
if new_c64 < c0_64 {
let (n_c2, ovf3) = rng.c[2].overflowing_add(1);
rng.c[2] = n_c2.into();
if ovf3 {
rng.c[3] += 1;
}
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_next_u32s(ptr: *mut Threefry32x4, out: *mut u32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
fry4_fill(buffer, rng.c, rng.k, rng.tw, |x| *x);
fry4_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_next_f32s(ptr: *mut Threefry32x4, out: *mut f32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
fry4_fill(buffer, rng.c, rng.k, rng.tw, |x| u2f_01!(f32, 32, *x));
fry4_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_rand_i32s(
ptr: *mut Threefry32x4,
out: *mut i32,
count: usize,
min: i32,
max: i32,
) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
let range = (max as i64 - min as i64 + 1) as u64;
fry4_fill(buffer, rng.c, rng.k, rng.tw, |x| {
*(((x.cast::<u64>() * range) >> 32).cast::<i32>() + min)
});
fry4_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x4_rand_f32s(
ptr: *mut Threefry32x4,
out: *mut f32,
count: usize,
min: f32,
max: f32,
) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
let mult = max - min;
fry4_fill(buffer, rng.c, rng.k, rng.tw, |x| {
u2f_01!(f32, 32, *x) * mult + min
});
fry4_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x2_new(seed: u32) -> *mut Threefry32x2 {
Box::into_raw(Box::new(Threefry32x2::new(seed)))
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x2_free(ptr: *mut Threefry32x2) {
if !ptr.is_null() {
unsafe {
let _ = Box::from_raw(ptr);
}
}
}
const THREEFRY32X2_PAR_CHUNK: usize = 0x20000;
#[inline(always)]
fn fry2_fill<T, M>(buffer: &mut [T], c0: [wu32; 2], k: [wu32; 3], map: M)
where
T: Copy + Default + Send,
M: Fn(wu32) -> T + Sync,
{
let nt = prefer_nt::<T>(buffer.len());
buffer
.par_chunks_mut(THREEFRY32X2_PAR_CHUNK)
.enumerate()
.for_each(|(chunk_idx, chunk)| {
let chunk_base = (chunk_idx * (THREEFRY32X2_PAR_CHUNK / 2)) as u64;
let c0_64 = (c0[0].cast::<u64>()) | ((c0[1].cast::<u64>()) << 32);
let mut c64 = c0_64 + chunk_base;
unsafe {
fill_chunk_auto(chunk, nt, || {
let mut out = [T::default(); 64];
for b in 0..32 {
let cc = c64 + (b as u64);
let c = [cc.cast(), (cc >> 32).cast()];
let r = Threefry32x2::compute(c, &k);
out[b << 1] = map(r[0]);
out[(b << 1) + 1] = map(r[1]);
}
c64 += 32;
out
});
}
});
}
#[inline(always)]
fn fry2_advance(rng: &mut Threefry32x2, count: usize) {
let num_blocks = (count.div_ceil(64) << 5) as u64;
let c0_64 = (rng.c[0].cast::<u64>()) | ((rng.c[1].cast::<u64>()) << 32);
let new_c64 = c0_64 + num_blocks;
rng.c[0] = new_c64.cast();
rng.c[1] = (new_c64 >> 32).cast();
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x2_next_u32s(ptr: *mut Threefry32x2, out: *mut u32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
fry2_fill(buffer, rng.c, rng.k, |x| *x);
fry2_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x2_next_f32s(ptr: *mut Threefry32x2, out: *mut f32, count: usize) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
fry2_fill(buffer, rng.c, rng.k, |x| u2f_01!(f32, 32, *x));
fry2_advance(rng, count);
}
}
#[unsafe(no_mangle)]
pub extern "C" fn threefry32x2_rand_i32s(
ptr: *mut Threefry32x2,
out: *mut i32,
count: usize,
min: i32,
max: i32,
) {
unsafe {
let rng = &mut *ptr;
let buffer = from_raw_parts_mut(out, count);
let range = (max as i64 - min as i64 + 1) as u64;
fry2_fill(buffer, rng.c, rng.k, |x| {
*((x.cast::<u64>() * range) >> 32).cast::<i32>() + min
});
fry2_advance(rng, count);
}
}