#![allow(unused, dead_code)]
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
#[cfg(all(feature = "simd", target_arch = "x86_64"))]
pub(crate) mod simd_f01 {
use std::arch::x86_64::*;
#[inline(always)]
pub(crate) unsafe fn u32x4(v: __m128i) -> __m128 {
unsafe {
let bits = _mm_or_si128(_mm_srli_epi32(v, 9), _mm_set1_epi32(0x3F80_0000));
_mm_sub_ps(_mm_castsi128_ps(bits), _mm_set1_ps(1.0))
}
}
#[inline]
#[target_feature(enable = "avx2")]
pub(crate) unsafe fn u32x8(v: __m256i) -> __m256 {
let bits = _mm256_or_si256(_mm256_srli_epi32(v, 9), _mm256_set1_epi32(0x3F80_0000));
_mm256_sub_ps(_mm256_castsi256_ps(bits), _mm256_set1_ps(1.0))
}
#[inline]
#[target_feature(enable = "avx512f")]
pub(crate) unsafe fn u32x16(v: __m512i) -> __m512 {
let bits = _mm512_or_si512(_mm512_srli_epi32(v, 9), _mm512_set1_epi32(0x3F80_0000));
_mm512_sub_ps(_mm512_castsi512_ps(bits), _mm512_set1_ps(1.0))
}
#[inline(always)]
pub(crate) unsafe fn u64x2(v: __m128i) -> __m128d {
unsafe {
let bits = _mm_or_si128(
_mm_srli_epi64(v, 11),
_mm_set1_epi64x(0x3FF0_0000_0000_0000),
);
_mm_sub_pd(_mm_castsi128_pd(bits), _mm_set1_pd(1.0))
}
}
#[inline]
#[target_feature(enable = "avx2")]
pub(crate) unsafe fn u64x4(v: __m256i) -> __m256d {
let bits = _mm256_or_si256(
_mm256_srli_epi64(v, 11),
_mm256_set1_epi64x(0x3FF0_0000_0000_0000),
);
_mm256_sub_pd(_mm256_castsi256_pd(bits), _mm256_set1_pd(1.0))
}
#[inline]
#[target_feature(enable = "avx512f")]
pub(crate) unsafe fn u64x8(v: __m512i) -> __m512d {
let bits = _mm512_or_si512(
_mm512_srli_epi64(v, 11),
_mm512_set1_epi64(0x3FF0_0000_0000_0000),
);
_mm512_sub_pd(_mm512_castsi512_pd(bits), _mm512_set1_pd(1.0))
}
}
macro_rules! randi_wide {
(i 32) => {
i64
};
(i 64) => {
i128
};
(u 32) => {
u64
};
(u 64) => {
u128
};
}
pub(crate) use randi_wide;
macro_rules! i2f_bits {
(32 bits) => {
0x3F800000
};
(32 bias) => {
9
};
(64 bits) => {
0x3FF0000000000000
};
(64 bias) => {
11
};
}
pub(crate) use i2f_bits;
macro_rules! u2f_01 {
($ft:ty, $bits:tt, $x:expr) => {{
<$ft>::from_bits(($x >> i2f_bits!($bits bias)) | i2f_bits!($bits bits)) - 1.0
}};
}
pub(crate) use u2f_01;
macro_rules! u2f_01w {
($ft:ty, $bits:tt, $x:expr) => {
::wrapn::wrap!(<$ft>::from_bits(($x.0.0 >> i2f_bits!($bits bias)) | i2f_bits!($bits bits)) - 1.0)
};
}
pub(crate) use u2f_01w;
macro_rules! sm64_from_seed32 {
($seed:expr) => {{
let mut s = $crate::SplitMix32::new($seed);
let sg =
$crate::SplitMix64::new(((s.nextu_const() as u64) << 32) | (s.nextu_const() as u64));
sg
}};
}
pub(crate) use sm64_from_seed32;
macro_rules! impl_ring_rng32 {
($ty:ty, $n:expr, $raw:ident) => {
impl $crate::rng::Rng for $ty {
type Word = u32;
#[inline]
fn nextu(&mut self) -> Self::Word {
if self.pos >= $n {
self.buf = self.$raw();
self.pos = ::wrapn::wrap!(0);
}
let v = self.buf[*self.pos];
self.pos += 1;
*v
}
}
};
}
pub(crate) use impl_ring_rng32;
macro_rules! impl_ring_rng64 {
($ty:ty, $n:expr, $raw:ident) => {
impl $crate::rng::Rng for $ty {
type Word = u64;
#[inline]
fn nextu(&mut self) -> Self::Word {
if self.pos >= $n {
self.buf = self.$raw();
self.pos = ::wrapn::wrap!(0);
}
let v = self.buf[*self.pos];
self.pos += 1;
*v
}
}
};
}
pub(crate) use impl_ring_rng64;
static DEFAULT_SEED_COUNTER: AtomicU64 = AtomicU64::new(0);
pub(crate) fn default_seed64() -> u64 {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
let count = DEFAULT_SEED_COUNTER.fetch_add(1, Ordering::Relaxed);
crate::prng::b64::SplitMix64::compute(nanos ^ count.wrapping_mul(0x9E3779B97F4A7C15))
}
pub(crate) fn default_seed32() -> u32 {
let z = default_seed64();
(z ^ (z >> 32)) as u32
}
#[inline(always)]
pub(crate) unsafe fn fill_with<T, F: FnMut() -> T>(out: *mut T, count: usize, mut next: F) {
let buffer = unsafe { std::slice::from_raw_parts_mut(out, count) };
for v in buffer {
*v = next();
}
}
#[inline(always)]
pub(crate) unsafe fn fill_chunk<T: Copy, const N: usize, F: FnMut() -> [T; N]>(
chunk: &mut [T],
mut generate: F,
) {
let mut out_ptr = chunk.as_mut_ptr();
let mut remaining = chunk.len();
while remaining >= N * 4 {
let v0 = generate();
let v1 = generate();
let v2 = generate();
let v3 = generate();
unsafe {
std::ptr::copy_nonoverlapping(v0.as_ptr(), out_ptr, N);
std::ptr::copy_nonoverlapping(v1.as_ptr(), out_ptr.add(N), N);
std::ptr::copy_nonoverlapping(v2.as_ptr(), out_ptr.add(N * 2), N);
std::ptr::copy_nonoverlapping(v3.as_ptr(), out_ptr.add(N * 3), N);
out_ptr = out_ptr.add(N * 4);
}
remaining -= N * 4;
}
while remaining >= N {
let v = generate();
unsafe {
std::ptr::copy_nonoverlapping(v.as_ptr(), out_ptr, N);
out_ptr = out_ptr.add(N);
}
remaining -= N;
}
if remaining > 0 {
let v = generate();
unsafe { std::ptr::copy_nonoverlapping(v.as_ptr(), out_ptr, remaining) };
}
}
#[cfg(all(feature = "simd", target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
#[allow(unsafe_op_in_unsafe_fn)]
pub(crate) unsafe fn fill_chunk_nt<T: Copy, const N: usize, F: FnMut() -> [T; N]>(
chunk: &mut [T],
mut generate: F,
) {
use std::arch::x86_64::*;
let words = (N * size_of::<T>()) / 32;
if words == 0 || !(N * size_of::<T>()).is_multiple_of(32) {
return fill_chunk(chunk, generate);
}
let mut p = chunk.as_mut_ptr();
let mut rem = chunk.len();
if (p as usize) & 31 == 0 {
while rem >= N {
let v = generate();
let src = v.as_ptr() as *const __m256i;
for i in 0..words {
_mm256_stream_si256(
(p as *mut u8).add(i * 32) as *mut __m256i,
_mm256_loadu_si256(src.add(i)),
);
}
p = p.add(N);
rem -= N;
}
_mm_sfence();
} else {
while rem >= N {
let v = generate();
std::ptr::copy_nonoverlapping(v.as_ptr(), p, N);
p = p.add(N);
rem -= N;
}
}
if rem > 0 {
let v = generate();
std::ptr::copy_nonoverlapping(v.as_ptr(), p, rem);
}
}
pub(crate) const NT_THRESHOLD_BYTES: usize = 24 << 20;
#[inline(always)]
pub(crate) fn prefer_nt<T>(total_elems: usize) -> bool {
total_elems * size_of::<T>() > NT_THRESHOLD_BYTES
}
#[inline(always)]
pub(crate) fn prefer_nt_for<T>(total_elems: usize, _sample: &[T]) -> bool {
total_elems * size_of::<T>() > NT_THRESHOLD_BYTES
}
#[inline(always)]
pub(crate) unsafe fn fill_chunk_auto<T: Copy, const N: usize, F: FnMut() -> [T; N]>(
chunk: &mut [T],
nt: bool,
generate: F,
) {
#[cfg(all(feature = "simd", target_arch = "x86_64"))]
if nt && std::arch::is_x86_feature_detected!("avx2") {
return unsafe { fill_chunk_nt(chunk, generate) };
}
let _ = nt;
unsafe { fill_chunk(chunk, generate) }
}
#[cfg(feature = "cabi")]
pub(crate) fn par_fill_reseed32<R, T, NF, SF>(
buffer: &mut [T],
base_seed: u32,
new_rng: NF,
step: SF,
) where
T: Copy + Default + Send,
NF: Fn(u32) -> R + Sync,
SF: Fn(&mut R) -> T + Sync,
{
use rayon::iter::{IndexedParallelIterator, ParallelIterator};
use rayon::slice::ParallelSliceMut;
const PAR_CHUNK: usize = 0x20000;
let nt = prefer_nt::<T>(buffer.len());
buffer
.par_chunks_mut(PAR_CHUNK)
.enumerate()
.for_each(|(chunk_idx, chunk)| {
let mut rng = new_rng(chunk_seed32(base_seed, chunk_idx));
unsafe {
fill_chunk_auto(chunk, nt, || {
let mut out = [T::default(); 16];
for v in &mut out {
*v = step(&mut rng);
}
out
});
}
});
}
#[cfg(feature = "cabi")]
pub(crate) fn par_fill_reseed64<R, T, NF, SF>(
buffer: &mut [T],
base_seed: u64,
new_rng: NF,
step: SF,
) where
T: Copy + Default + Send,
NF: Fn(u64) -> R + Sync,
SF: Fn(&mut R) -> T + Sync,
{
use rayon::iter::{IndexedParallelIterator, ParallelIterator};
use rayon::slice::ParallelSliceMut;
const PAR_CHUNK: usize = 0x20000;
let nt = prefer_nt::<T>(buffer.len());
buffer
.par_chunks_mut(PAR_CHUNK)
.enumerate()
.for_each(|(chunk_idx, chunk)| {
let chunk_seed = crate::prng::b64::SplitMix64::compute(
base_seed.wrapping_add((chunk_idx as u64).wrapping_mul(0x9E3779B97F4A7C15)),
);
let mut rng = new_rng(chunk_seed);
unsafe {
fill_chunk_auto(chunk, nt, || {
let mut out = [T::default(); 8];
for v in &mut out {
*v = step(&mut rng);
}
out
});
}
});
}
#[inline]
pub(crate) fn chunk_seed32(base_seed: u32, chunk_idx: usize) -> u32 {
let x = base_seed.wrapping_add((chunk_idx as u32).wrapping_mul(0x9E37_79B9));
let mut z = x as u64;
z ^= z >> 16;
z = z.wrapping_mul(0xFF51_AFD7_ED55_8CCD);
z ^= z >> 16;
z = z.wrapping_mul(0xC4CE_B9FE_1A85_EC53);
(z ^ (z >> 16)) as u32
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_seed_calls_are_distinct() {
assert_ne!(default_seed64(), default_seed64());
assert_ne!(default_seed32(), default_seed32());
}
fn counter_batches<const N: usize>() -> impl FnMut() -> [u32; N] {
let mut next = 0u32;
move || {
let mut out = [0u32; N];
for v in &mut out {
*v = next;
next += 1;
}
out
}
}
fn check_fill(buf: &[u32], len: usize) {
for (i, &v) in buf[..len].iter().enumerate() {
assert_eq!(v as usize, i, "element {i} wrong");
}
}
#[test]
fn fill_chunk_writes_every_element() {
for len in [0usize, 1, 7, 16, 63, 64, 65, 1000] {
let mut buf = vec![u32::MAX; len];
unsafe { fill_chunk::<u32, 16, _>(&mut buf, counter_batches()) };
check_fill(&buf, len);
}
}
#[cfg(all(feature = "simd", target_arch = "x86_64"))]
#[test]
fn fill_chunk_nt_writes_every_element() {
if !std::arch::is_x86_feature_detected!("avx2") {
return;
}
for len in [0usize, 1, 7, 16, 63, 64, 65, 1000] {
let mut buf = vec![u32::MAX; len];
unsafe { fill_chunk_nt::<u32, 16, _>(&mut buf, counter_batches()) };
check_fill(&buf, len);
}
}
#[cfg(all(feature = "simd", target_arch = "x86_64"))]
#[test]
fn fill_chunk_nt_small_batch_falls_back() {
if !std::arch::is_x86_feature_detected!("avx2") {
return;
}
let mut buf = vec![u32::MAX; 100];
unsafe { fill_chunk_nt::<u32, 4, _>(&mut buf, counter_batches()) };
check_fill(&buf, 100);
}
}
#[cfg(feature = "simd")]
macro_rules! dispatch_simd {
($ret_type:ty, $fallback_fn:ident, $avx512_fn:ident, $seed:expr) => {{
#[cfg(target_arch = "x86_64")]
if std::arch::is_x86_feature_detected!("avx512f") {
return $avx512_fn($seed) as *mut $ret_type;
}
$fallback_fn($seed) as *mut $ret_type
}};
($avx512_type:ty, $fallback_type:ty, $fallback_fn:ident, $avx512_fn:ident, $ptr:expr $(, $arg:expr)*) => {{
#[cfg(target_arch = "x86_64")]
if std::arch::is_x86_feature_detected!("avx512f") {
$avx512_fn($ptr as *mut $avx512_type $(, $arg)*);
return;
}
$fallback_fn($ptr as *mut $fallback_type $(, $arg)*);
}};
}
#[cfg(feature = "simd")]
pub(crate) use dispatch_simd;
macro_rules! safe_test {
($($name:ident),+) => {
$(pastey::paste! {
#[test]
fn [<test_ $name:snake>]() {
let mut rng1 = $name::new(0);
let mut rng2 = $name::new(0);
assert_eq!(rng1.nextu(), rng2.nextu());
assert_eq!(rng1.nextf(), rng2.nextf());
}
})+
};
}
pub(crate) use safe_test;
macro_rules! unsafe_test {
($($name:ident),+) => {
$(pastey::paste! {
#[test]
fn [<test_ $name:snake>]() {
unsafe fn as_bytes<T>(v: &T) -> &[u8] {
unsafe {
std::slice::from_raw_parts(v as *const T as *const u8, std::mem::size_of::<T>())
}
}
unsafe {
let mut rng1 = $name::new(0);
let mut rng2 = $name::new(0);
let u1 = rng1.nextuv();
let u2 = rng2.nextuv();
assert_eq!(as_bytes(&u1), as_bytes(&u2));
let f1 = rng1.nextfv();
let f2 = rng2.nextfv();
assert_eq!(as_bytes(&f1), as_bytes(&f2));
}
}
})+
};
}
pub(crate) use unsafe_test;
macro_rules! impl_default_from_seed32 {
($($type:ty),* $(,)?) => {
$(
impl Default for $type {
#[inline]
fn default() -> Self {
Self::new($crate::_internal::default_seed32())
}
}
impl $crate::Seed for $type {
type Seed = u32;
#[inline]
fn from_seed(seed: u32) -> Self {
Self::new(seed)
}
}
)*
};
}
pub(crate) use impl_default_from_seed32;
macro_rules! impl_default_from_seed64 {
($($type:ty),* $(,)?) => {
$(
impl Default for $type {
#[inline]
fn default() -> Self {
Self::new($crate::_internal::default_seed64())
}
}
impl $crate::Seed for $type {
type Seed = u64;
#[inline]
fn from_seed(seed: u64) -> Self {
Self::new(seed)
}
}
)*
};
}
pub(crate) use impl_default_from_seed64;
#[cfg(feature = "rand")]
macro_rules! impl_try_rng_trait {
($($type:ty),* $(,)?) => {
$(
impl rand_core::TryRng for $type {
type Error = std::convert::Infallible;
fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
Ok(self.nextu())
}
fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
let hi = self.nextu() as u64;
let lo = self.nextu() as u64;
Ok(hi << 32 | lo)
}
fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> {
let mut i = 0;
while i < dst.len() {
let remaining = dst.len() - i;
if remaining >= 4 {
let val = self.nextu();
dst[i..i + 4].copy_from_slice(&val.to_le_bytes());
i += 4;
} else {
let val = self.nextu();
dst[i..].copy_from_slice(&val.to_le_bytes()[..remaining]);
i += remaining;
}
}
Ok(())
}
}
)*
};
}
#[cfg(feature = "rand")]
pub(crate) use impl_try_rng_trait;
#[cfg(feature = "rand")]
macro_rules! impl_rand_trait {
($($type:ty),* $(,)?) => {
$(
impl rand_core::SeedableRng for $type {
type Seed = [u8; 4];
fn from_seed(seed: Self::Seed) -> Self {
let seed = u32::from_ne_bytes(seed);
Self::new(seed.into())
}
}
)*
};
}
#[cfg(feature = "rand")]
pub(crate) use impl_rand_trait;