Skip to content

Commit b5f3d00

Browse files
committed
feat: update distr and add more simd
1 parent 2f27839 commit b5f3d00

17 files changed

Lines changed: 277 additions & 6 deletions

benches/dist_multicore.rs

Lines changed: 91 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,14 @@
11
use std::hint::black_box;
22
use std::time::Instant;
33

4+
use rand_distr::Distribution;
45
use rayon::ThreadPool;
56
use rayon::ThreadPoolBuilder;
67
use stochastic_rs::distributions::DistributionSampler;
78
use stochastic_rs::distributions::exp::SimdExp;
89
use stochastic_rs::distributions::normal::SimdNormal;
910
use stochastic_rs::distributions::poisson::SimdPoisson;
11+
use stochastic_rs::simd_rng::SimdRng;
1012

1113
fn median_ms(samples: &mut [f64]) -> f64 {
1214
samples.sort_by(f64::total_cmp);
@@ -46,6 +48,78 @@ where
4648
median_ms(&mut times_ms)
4749
}
4850

51+
fn bench_normal_fill_slice(n: usize, warmup: usize, runs: usize) -> (f64, f64, f64) {
52+
let simd = SimdNormal::<f64>::new(0.0, 1.0);
53+
let rand_distr = rand_distr::Normal::<f64>::new(0.0, 1.0).expect("valid normal params");
54+
let mut out = vec![0.0f64; n];
55+
let iters = (262_144 / n.max(1)).clamp(1, 16_384);
56+
57+
let mut simd_rng = SimdRng::new();
58+
for _ in 0..warmup {
59+
for _ in 0..iters {
60+
simd.fill_slice(&mut simd_rng, &mut out);
61+
}
62+
black_box(&out);
63+
}
64+
let mut simd_times_ms = Vec::with_capacity(runs);
65+
for _ in 0..runs {
66+
let t0 = Instant::now();
67+
for _ in 0..iters {
68+
simd.fill_slice(&mut simd_rng, &mut out);
69+
}
70+
black_box(&out);
71+
simd_times_ms.push(t0.elapsed().as_secs_f64() * 1_000.0 / iters as f64);
72+
}
73+
74+
let mut base_rng = SimdRng::new();
75+
for _ in 0..warmup {
76+
for _ in 0..iters {
77+
for x in &mut out {
78+
*x = rand_distr.sample(&mut base_rng);
79+
}
80+
}
81+
black_box(&out);
82+
}
83+
let mut base_times_ms = Vec::with_capacity(runs);
84+
for _ in 0..runs {
85+
let t0 = Instant::now();
86+
for _ in 0..iters {
87+
for x in &mut out {
88+
*x = rand_distr.sample(&mut base_rng);
89+
}
90+
}
91+
black_box(&out);
92+
base_times_ms.push(t0.elapsed().as_secs_f64() * 1_000.0 / iters as f64);
93+
}
94+
95+
let mut base_thread_rng = rand::rng();
96+
for _ in 0..warmup {
97+
for _ in 0..iters {
98+
for x in &mut out {
99+
*x = rand_distr.sample(&mut base_thread_rng);
100+
}
101+
}
102+
black_box(&out);
103+
}
104+
let mut base_thread_times_ms = Vec::with_capacity(runs);
105+
for _ in 0..runs {
106+
let t0 = Instant::now();
107+
for _ in 0..iters {
108+
for x in &mut out {
109+
*x = rand_distr.sample(&mut base_thread_rng);
110+
}
111+
}
112+
black_box(&out);
113+
base_thread_times_ms.push(t0.elapsed().as_secs_f64() * 1_000.0 / iters as f64);
114+
}
115+
116+
(
117+
median_ms(&mut simd_times_ms),
118+
median_ms(&mut base_times_ms),
119+
median_ms(&mut base_thread_times_ms),
120+
)
121+
}
122+
49123
fn run_case<T, D>(name: &str, dist: &D, m: usize, n: usize, single: &ThreadPool, multi: &ThreadPool)
50124
where
51125
D: DistributionSampler<T> + Clone + Send,
@@ -104,4 +178,21 @@ fn main() {
104178
&single,
105179
&multi,
106180
);
181+
182+
println!();
183+
println!("Normal fill_slice benchmark (single-thread)");
184+
println!("Reference A: rand_distr + SimdRng (fair algorithm compare)");
185+
println!("Reference B: rand_distr + rand::rng() (out-of-box baseline)");
186+
for &n in &[4usize, 8, 16, 64, 256, 4096, 65_536] {
187+
let (simd_ms, base_simd_ms, base_thread_ms) = bench_normal_fill_slice(n, 2, 9);
188+
let speedup_a = base_simd_ms / simd_ms;
189+
let speedup_b = base_thread_ms / simd_ms;
190+
let simd_us = simd_ms * 1_000.0;
191+
let base_simd_us = base_simd_ms * 1_000.0;
192+
let base_thread_us = base_thread_ms * 1_000.0;
193+
println!(
194+
"{:>12} | n={n:<6} | simd={simd_us:>9.3} us | rd+simd_rng={base_simd_us:>9.3} us ({speedup_a:>5.2}x) | rd+rand_rng={base_thread_us:>9.3} us ({speedup_b:>5.2}x)",
195+
"Normal<f64>"
196+
);
197+
}
107198
}

src/distributions.rs

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,11 @@ impl SimdFloatExt for f32 {
288288
rng.random()
289289
}
290290

291+
#[inline(always)]
292+
fn sample_uniform_simd(rng: &mut crate::simd_rng::SimdRng) -> f32 {
293+
rng.next_f32()
294+
}
295+
291296
fn simd_from_i32x8(v: wide::i32x8) -> f32x8 {
292297
v.round_float()
293298
}
@@ -388,6 +393,11 @@ impl SimdFloatExt for f64 {
388393
rng.random()
389394
}
390395

396+
#[inline(always)]
397+
fn sample_uniform_simd(rng: &mut crate::simd_rng::SimdRng) -> f64 {
398+
rng.next_f64()
399+
}
400+
391401
fn simd_from_i32x8(v: wide::i32x8) -> f64x8 {
392402
f64x8::from_i32x8(v)
393403
}

src/distributions/alpha_stable.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -172,8 +172,8 @@ impl<T: SimdFloatExt> SimdAlphaStable<T> {
172172
let scale = self.scale;
173173
let loc = self.location;
174174
for x in out.iter_mut() {
175-
let mut u = T::sample_uniform(rng);
176-
let mut e = T::sample_uniform(rng);
175+
let mut u = T::sample_uniform_simd(rng);
176+
let mut e = T::sample_uniform_simd(rng);
177177
u = Self::clamp_open_unit(u);
178178
e = Self::clamp_open_unit(e);
179179
let v = pi * (u - T::from(0.5).unwrap());

src/distributions/beta.rs

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@ use rand_distr::Distribution;
1212
use super::SimdFloatExt;
1313
use super::gamma::SimdGamma;
1414

15+
const SMALL_BETA_THRESHOLD: usize = 16;
16+
1517
pub struct SimdBeta<T: SimdFloatExt> {
1618
alpha: T,
1719
beta: T,
@@ -39,6 +41,15 @@ impl<T: SimdFloatExt> SimdBeta<T> {
3941
}
4042

4143
pub fn fill_slice_fast(&self, out: &mut [T]) {
44+
if out.len() < SMALL_BETA_THRESHOLD {
45+
let mut rng = crate::simd_rng::SimdRng::new();
46+
for x in out.iter_mut() {
47+
let a = self.gamma1.sample(&mut rng);
48+
let b = self.gamma2.sample(&mut rng);
49+
*x = a / (a + b);
50+
}
51+
return;
52+
}
4253
let mut g1 = [T::zero(); 8];
4354
let mut g2 = [T::zero(); 8];
4455
let mut chunks = out.chunks_exact_mut(8);

src/distributions/cauchy.rs

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@ use rand_distr::Distribution;
1212
use super::SimdFloatExt;
1313
use crate::simd_rng::SimdRng;
1414

15+
const SMALL_CAUCHY_THRESHOLD: usize = 16;
16+
1517
pub struct SimdCauchy<T: SimdFloatExt> {
1618
x0: T,
1719
gamma: T,
@@ -38,6 +40,15 @@ impl<T: SimdFloatExt> SimdCauchy<T> {
3840

3941
pub fn fill_slice_fast(&self, out: &mut [T]) {
4042
let rng = unsafe { &mut *self.simd_rng.get() };
43+
if out.len() < SMALL_CAUCHY_THRESHOLD {
44+
let pi = T::pi();
45+
let half = T::from(0.5).unwrap();
46+
for x in out.iter_mut() {
47+
let u = T::sample_uniform_simd(rng);
48+
*x = self.x0 + self.gamma * (pi * (u - half)).tan();
49+
}
50+
return;
51+
}
4152
let x0 = T::splat(self.x0);
4253
let g = T::splat(self.gamma);
4354
let pi = T::splat(T::pi());

src/distributions/exp.rs

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@ use crate::simd_rng::SimdRng;
1818
const ZIG_EXP_R: f64 = 7.697_117_470_131_487;
1919
const ZIG_EXP_V: f64 = 3.949_659_822_581_572e-3;
2020
const TABLE_SIZE: usize = 256;
21+
const SMALL_EXP_THRESHOLD: usize = 16;
2122

2223
struct ExpZigTables {
2324
ke: [i32; TABLE_SIZE],
@@ -115,10 +116,28 @@ impl<T: SimdFloatExt, const N: usize> SimdExpZig<T, N> {
115116
}
116117
}
117118

119+
#[inline]
120+
fn sample_exp1_one(rng: &mut SimdRng, tables: &ExpZigTables) -> T {
121+
let hz = rng.next_i32();
122+
let iz = (hz & 0xFF) as usize;
123+
let abs_hz = hz.unsigned_abs() as i64;
124+
if abs_hz < tables.ke[iz] as i64 {
125+
T::from_f64_fast((abs_hz as f64) * tables.we[iz])
126+
} else {
127+
efix::<T>(hz, iz, tables, rng)
128+
}
129+
}
130+
118131
#[inline]
119132
fn fill_exp1(buf: &mut [T], rng: &mut SimdRng) {
120133
let tables = exp_zig_tables();
121134
let len = buf.len();
135+
if len < SMALL_EXP_THRESHOLD {
136+
for x in buf.iter_mut() {
137+
*x = Self::sample_exp1_one(rng, tables);
138+
}
139+
return;
140+
}
122141
let mask255 = i32x8::splat(0xFF);
123142
let mut filled = 0;
124143

src/distributions/gamma.rs

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,7 @@ impl<T: SimdFloatExt> SimdGamma<T> {
5959
if v <= T::zero() {
6060
continue;
6161
}
62-
let u: T = T::sample_uniform(rng);
62+
let u: T = T::sample_uniform_simd(rng);
6363
let z2 = z * z;
6464
if u < T::one() - c1 * z2 * z2 {
6565
break d * v;
@@ -68,7 +68,7 @@ impl<T: SimdFloatExt> SimdGamma<T> {
6868
break d * v;
6969
}
7070
};
71-
let u: T = T::sample_uniform(rng);
71+
let u: T = T::sample_uniform_simd(rng);
7272
*x = self.scale * g * u.powf(inv_alpha);
7373
}
7474
} else {
@@ -81,7 +81,7 @@ impl<T: SimdFloatExt> SimdGamma<T> {
8181
if v <= T::zero() {
8282
continue;
8383
}
84-
let u: T = T::sample_uniform(rng);
84+
let u: T = T::sample_uniform_simd(rng);
8585
let z2 = z * z;
8686
if u < T::one() - c1 * z2 * z2 {
8787
break d * v;

src/distributions/geometric.rs

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,8 @@ use wide::f64x8;
1313

1414
use crate::simd_rng::SimdRng;
1515

16+
const SMALL_GEOMETRIC_THRESHOLD: usize = 16;
17+
1618
pub struct SimdGeometric<T: PrimInt> {
1719
p: f64,
1820
buffer: UnsafeCell<[T; 16]>,
@@ -37,6 +39,15 @@ impl<T: PrimInt> SimdGeometric<T> {
3739
pub fn fill_slice_fast(&self, out: &mut [T]) {
3840
let rng = unsafe { &mut *self.simd_rng.get() };
3941
let ln1p = (1.0 - self.p).ln();
42+
if out.len() < SMALL_GEOMETRIC_THRESHOLD {
43+
let inv_ln1p = 1.0 / ln1p;
44+
for x in out.iter_mut() {
45+
let u = rng.next_f64();
46+
let g = (u.ln() * inv_ln1p).floor();
47+
*x = num_traits::cast(g.max(0.0) as u64).unwrap_or(T::zero());
48+
}
49+
return;
50+
}
4051
let inv_ln1p = f64x8::splat(1.0 / ln1p);
4152
let mut chunks = out.chunks_exact_mut(8);
4253
for chunk in &mut chunks {

src/distributions/inverse_gauss.rs

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,8 @@ use super::SimdFloatExt;
1313
use super::normal::SimdNormal;
1414
use crate::simd_rng::SimdRng;
1515

16+
const SMALL_INVERSE_GAUSS_THRESHOLD: usize = 16;
17+
1618
pub struct SimdInverseGauss<T: SimdFloatExt> {
1719
mu: T,
1820
lambda: T,
@@ -41,6 +43,25 @@ impl<T: SimdFloatExt> SimdInverseGauss<T> {
4143

4244
pub fn fill_slice_fast(&self, out: &mut [T]) {
4345
let rng = unsafe { &mut *self.simd_rng.get() };
46+
if out.len() < SMALL_INVERSE_GAUSS_THRESHOLD {
47+
let two = T::from(2.0).unwrap();
48+
let four = T::from(4.0).unwrap();
49+
for x in out.iter_mut() {
50+
let z = self.normal.sample(rng);
51+
let u = T::sample_uniform_simd(rng);
52+
let w = z * z;
53+
let t1 = self.mu + (self.mu * self.mu * w) / (two * self.lambda);
54+
let rad = (four * self.mu * self.lambda * w + self.mu * self.mu * w * w).sqrt();
55+
let xr = t1 - (self.mu / (two * self.lambda)) * rad;
56+
let check = self.mu / (self.mu + xr);
57+
*x = if u < check {
58+
xr
59+
} else {
60+
self.mu * self.mu / xr
61+
};
62+
}
63+
return;
64+
}
4465
let two = T::splat(T::from(2.0).unwrap());
4566
let four = T::splat(T::from(4.0).unwrap());
4667
let mu = T::splat(self.mu);

src/distributions/normal.rs

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ struct ZigTables {
2323
}
2424

2525
static ZIG_TABLES: OnceLock<ZigTables> = OnceLock::new();
26+
const SMALL_NORMAL_THRESHOLD: usize = 16;
2627

2728
fn zig_tables() -> &'static ZigTables {
2829
ZIG_TABLES.get_or_init(|| {
@@ -174,10 +175,28 @@ impl<T: SimdFloatExt, const N: usize> SimdNormal<T, N> {
174175
}
175176
}
176177

178+
#[inline]
179+
fn sample_one(rng: &mut SimdRng, tables: &ZigTables, mean: T, std_dev: T) -> T {
180+
let hz = rng.next_i32();
181+
let iz = (hz & 127) as usize;
182+
let z = if (hz.unsigned_abs() as i64) < tables.kn[iz] as i64 {
183+
T::from_f64_fast(hz as f64 * tables.wn[iz])
184+
} else {
185+
nfix::<T>(hz, iz, tables, rng)
186+
};
187+
mean + std_dev * z
188+
}
189+
177190
#[inline]
178191
fn fill_ziggurat(buf: &mut [T], rng: &mut SimdRng, mean: T, std_dev: T) {
179192
let len = buf.len();
180193
let tables = zig_tables();
194+
if len < SMALL_NORMAL_THRESHOLD {
195+
for x in buf.iter_mut() {
196+
*x = Self::sample_one(rng, tables, mean, std_dev);
197+
}
198+
return;
199+
}
181200
let mean_simd = T::splat(mean);
182201
let std_dev_simd = T::splat(std_dev);
183202
let mask127 = i32x8::splat(127);
@@ -306,10 +325,27 @@ impl<T: SimdFloatExt, const N: usize> SimdNormal<T, N> {
306325
(a, b)
307326
}
308327

328+
#[inline]
329+
fn sample_one_standard(rng: &mut SimdRng, tables: &ZigTables) -> T {
330+
let hz = rng.next_i32();
331+
let iz = (hz & 127) as usize;
332+
if (hz.unsigned_abs() as i64) < tables.kn[iz] as i64 {
333+
T::from_f64_fast(hz as f64 * tables.wn[iz])
334+
} else {
335+
nfix::<T>(hz, iz, tables, rng)
336+
}
337+
}
338+
309339
#[inline]
310340
fn fill_ziggurat_standard(buf: &mut [T], rng: &mut SimdRng) {
311341
let len = buf.len();
312342
let tables = zig_tables();
343+
if len < SMALL_NORMAL_THRESHOLD {
344+
for x in buf.iter_mut() {
345+
*x = Self::sample_one_standard(rng, tables);
346+
}
347+
return;
348+
}
313349
let mask127 = i32x8::splat(127);
314350
let mut filled = 0;
315351

0 commit comments

Comments
 (0)