Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_core/src/rand.rs

18.2 KiB, 127 runs

created by r1870400018:118, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1use crate::{
2 prelude::*,
3 byte::B32,
4};
5
6use std::cmp::PartialOrd;
7
8use rand::{
9 thread_rng,
10 Rng,
11};
12use rand_core::{
13 OsRng,
14 RngCore,
15};
16
17
18/// Sampling method for range generation.
19#[derive(Clone, Copy, Debug)]
20pub enum SamplingMethod {
21 Uniform,
22 GaussianClampedDerived,
23 GaussianClampedExplicit { mean: f32, stdev: f32 },
24}
25
26pub trait RanDef {
27 fn randef() -> Self where Self: Sized;
28}
29
30impl RanDef for u8 {
31 fn randef() -> Self { Rand::rand_u8() }
32}
33
34impl RanDef for u16 {
35 fn randef() -> Self { Rand::rand_u16() }
36}
37
38impl RanDef for u32 {
39 fn randef() -> Self { Rand::rand_u32() }
40}
41
42impl RanDef for u64 {
43 fn randef() -> Self { Rand::rand_u64() }
44}
45
46impl RanDef for u128 {
47 fn randef() -> Self { Rand::rand_u128() }
48}
49
50impl RanDef for B32 {
51 fn randef() -> Self {
52 let mut a = [0; 32];
53 Rand::fill_u8(&mut a);
54 Self(a)
55 }
56}
57
58pub struct Rand;
59
60impl Rand {
61 pub fn generate_random_string(
62 len: usize,
63 charset: &str,
64 )
65 -> String
66 {
67 let charset = charset.as_bytes();
68 let mut rng = thread_rng();
69 let pass: String = (0..len)
70 .map(|_| {
71 let idx = rng.gen_range(0..charset.len());
72 charset[idx] as char
73 })
74 .collect();
75 pass
76 }
77
78 pub fn value<T>() -> T
79 where
80 rand::distributions::Standard: rand::distributions::Distribution<T>,
81 {
82 let mut rng = thread_rng();
83 rng.gen()
84 }
85
86 pub fn in_range<T>(
87 lower: T,
88 upper: T,
89 )
90 -> T
91 where
92 T: PartialOrd + rand::distributions::uniform::SampleUniform
93 {
94 let mut rng = thread_rng();
95 rng.gen_range(lower..=upper)
96 }
97
98 pub fn rand_u8() -> u8 {
99 OsRng.next_u32() as u8
100 }
101
102 pub fn rand_u16() -> u16 {
103 OsRng.next_u32() as u16
104 }
105
106 pub fn rand_u32() -> u32 {
107 OsRng.next_u32()
108 }
109
110 pub fn rand_u64() -> u64 {
111 OsRng.next_u64()
112 }
113
114 pub fn rand_u128() -> u128 {
115 let a = OsRng.next_u64() as u128;
116 let b = (OsRng.next_u64() as u128) << 64;
117 a | b
118 }
119
120 pub fn fill_u8(a: &mut [u8]) {
121 thread_rng().fill(&mut a[..]);
122 }
123
124 pub fn normal(
125 mean: f32,
126 stdev: f32,
127 )
128 -> f32
129 {
130 let mut rng = thread_rng();
131 let u1: f32 = loop {
132 let val = rng.gen::<f32>();
133 if val > 0.0 {
134 break val;
135 }
136 };
137 let u2: f32 = rng.gen();
138 let z0 = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f32::consts::PI * u2).cos();
139 mean + stdev * z0
140 }
141
142 pub fn normal_f64(
143 mean: f64,
144 stdev: f64,
145 )
146 -> f64
147 {
148 let mut rng = thread_rng();
149 let u1: f64 = loop {
150 let val = rng.gen::<f64>();
151 if val > 0.0 {
152 break val;
153 }
154 };
155 let u2: f64 = rng.gen();
156 let z0 = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
157 mean + stdev * z0
158 }
159
160 // Sampling methods for all numeric types.
161
162 pub fn sample_f32(
163 min: f32,
164 max: f32,
165 method: SamplingMethod,
166 )
167 -> Outcome<f32>
168 {
169 if min > max {
170 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
171 }
172
173 if !min.is_finite() || !max.is_finite() {
174 return Err(err!("Range bounds must be finite: [{}, {}]", min, max; Invalid, Range));
175 }
176
177 Ok(match method {
178 SamplingMethod::Uniform => {
179 let scale = max - min;
180 min + scale * Self::value::<f32>()
181 }
182 SamplingMethod::GaussianClampedDerived => {
183 let mean = (min + max) / 2.0;
184 let stdev = (max - min) / 6.0;
185 let sample = Self::normal(mean, stdev);
186 sample.max(min).min(max)
187 }
188 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
189 let sample = Self::normal(mean, stdev);
190 sample.max(min).min(max)
191 }
192 })
193 }
194
195 pub fn sample_f64(
196 min: f64,
197 max: f64,
198 method: SamplingMethod,
199 )
200 -> Outcome<f64>
201 {
202 if min > max {
203 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
204 }
205
206 if !min.is_finite() || !max.is_finite() {
207 return Err(err!("Range bounds must be finite: [{}, {}]", min, max; Invalid, Range));
208 }
209
210 Ok(match method {
211 SamplingMethod::Uniform => {
212 let scale = max - min;
213 min + scale * Self::value::<f64>()
214 }
215 SamplingMethod::GaussianClampedDerived => {
216 let mean = (min + max) / 2.0;
217 let stdev = (max - min) / 6.0;
218 let sample = Self::normal_f64(mean, stdev);
219 sample.max(min).min(max)
220 }
221 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
222 let sample = Self::normal_f64(mean as f64, stdev as f64);
223 sample.max(min).min(max)
224 }
225 })
226 }
227
228 pub fn sample_u8(
229 min: u8,
230 max: u8,
231 method: SamplingMethod,
232 )
233 -> Outcome<u8>
234 {
235 if min > max {
236 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
237 }
238
239 Ok(match method {
240 SamplingMethod::Uniform => Self::in_range(min, max),
241 SamplingMethod::GaussianClampedDerived => {
242 let mean = (min as f32 + max as f32) / 2.0;
243 let stdev = (max as f32 - min as f32) / 6.0; // 6σ covers ~99.7%.
244 let sample = Self::normal(mean, stdev).round();
245 sample.max(min as f32).min(max as f32) as u8
246 }
247 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
248 let sample = Self::normal(mean, stdev).round();
249 sample.max(min as f32).min(max as f32) as u8
250 }
251 })
252 }
253
254 pub fn sample_u16(
255 min: u16,
256 max: u16,
257 method: SamplingMethod,
258 )
259 -> Outcome<u16>
260 {
261 if min > max {
262 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
263 }
264
265 Ok(match method {
266 SamplingMethod::Uniform => Self::in_range(min, max),
267 SamplingMethod::GaussianClampedDerived => {
268 let mean = (min as f32 + max as f32) / 2.0;
269 let stdev = (max as f32 - min as f32) / 6.0;
270 let sample = Self::normal(mean, stdev).round();
271 sample.max(min as f32).min(max as f32) as u16
272 }
273 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
274 let sample = Self::normal(mean, stdev).round();
275 sample.max(min as f32).min(max as f32) as u16
276 }
277 })
278 }
279
280 pub fn sample_u32(
281 min: u32,
282 max: u32,
283 method: SamplingMethod,
284 )
285 -> Outcome<u32>
286 {
287 if min > max {
288 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
289 }
290
291 Ok(match method {
292 SamplingMethod::Uniform => Self::in_range(min, max),
293 SamplingMethod::GaussianClampedDerived => {
294 let mean = (min as f64 + max as f64) / 2.0; // Use f64 for precision.
295 let stdev = (max as f64 - min as f64) / 6.0;
296 let sample = Self::normal_f64(mean, stdev).round();
297 sample.max(min as f64).min(max as f64) as u32
298 }
299 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
300 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
301 sample.max(min as f64).min(max as f64) as u32
302 }
303 })
304 }
305
306 pub fn sample_u64(
307 min: u64,
308 max: u64,
309 method: SamplingMethod,
310 )
311 -> Outcome<u64>
312 {
313 if min > max {
314 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
315 }
316
317 Ok(match method {
318 SamplingMethod::Uniform => Self::in_range(min, max),
319 SamplingMethod::GaussianClampedDerived => {
320 let mean = (min as f64 + max as f64) / 2.0;
321 let stdev = (max as f64 - min as f64) / 6.0;
322 let sample = Self::normal_f64(mean, stdev).round();
323 sample.max(min as f64).min(max as f64) as u64
324 }
325 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
326 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
327 sample.max(min as f64).min(max as f64) as u64
328 }
329 })
330 }
331
332 pub fn sample_u128(
333 min: u128,
334 max: u128,
335 method: SamplingMethod,
336 )
337 -> Outcome<u128>
338 {
339 if min > max {
340 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
341 }
342
343 Ok(match method {
344 SamplingMethod::Uniform => Self::in_range(min, max),
345 SamplingMethod::GaussianClampedDerived => {
346 // Careful handling for large values.
347 let min_f64 = min as f64;
348 let max_f64 = max as f64;
349 let mean = (min_f64 + max_f64) / 2.0;
350 let stdev = (max_f64 - min_f64) / 6.0;
351 let sample = Self::normal_f64(mean, stdev).round();
352
353 if sample < 0.0 {
354 min
355 } else if sample > u128::MAX as f64 {
356 max
357 } else {
358 sample.max(min_f64).min(max_f64) as u128
359 }
360 }
361 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
362 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
363 if sample < 0.0 {
364 min
365 } else if sample > u128::MAX as f64 {
366 max
367 } else {
368 sample.max(min as f64).min(max as f64) as u128
369 }
370 }
371 })
372 }
373
374 pub fn sample_usize(
375 min: usize,
376 max: usize,
377 method: SamplingMethod,
378 )
379 -> Outcome<usize>
380 {
381 if min > max {
382 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
383 }
384
385 Ok(match method {
386 SamplingMethod::Uniform => Self::in_range(min, max),
387 SamplingMethod::GaussianClampedDerived => {
388 let mean = (min as f64 + max as f64) / 2.0;
389 let stdev = (max as f64 - min as f64) / 6.0;
390 let sample = Self::normal_f64(mean, stdev).round();
391 sample.max(min as f64).min(max as f64) as usize
392 }
393 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
394 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
395 sample.max(min as f64).min(max as f64) as usize
396 }
397 })
398 }
399
400 pub fn sample_i8(
401 min: i8,
402 max: i8,
403 method: SamplingMethod,
404 )
405 -> Outcome<i8>
406 {
407 if min > max {
408 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
409 }
410
411 Ok(match method {
412 SamplingMethod::Uniform => Self::in_range(min, max),
413 SamplingMethod::GaussianClampedDerived => {
414 let mean = (min as f32 + max as f32) / 2.0;
415 let stdev = (max as f32 - min as f32) / 6.0;
416 let sample = Self::normal(mean, stdev).round();
417 sample.max(min as f32).min(max as f32) as i8
418 }
419 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
420 let sample = Self::normal(mean, stdev).round();
421 sample.max(min as f32).min(max as f32) as i8
422 }
423 })
424 }
425
426 pub fn sample_i16(
427 min: i16,
428 max: i16,
429 method: SamplingMethod,
430 )
431 -> Outcome<i16>
432 {
433 if min > max {
434 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
435 }
436
437 Ok(match method {
438 SamplingMethod::Uniform => Self::in_range(min, max),
439 SamplingMethod::GaussianClampedDerived => {
440 let mean = (min as f32 + max as f32) / 2.0;
441 let stdev = (max as f32 - min as f32) / 6.0;
442 let sample = Self::normal(mean, stdev).round();
443 sample.max(min as f32).min(max as f32) as i16
444 }
445 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
446 let sample = Self::normal(mean, stdev).round();
447 sample.max(min as f32).min(max as f32) as i16
448 }
449 })
450 }
451
452 pub fn sample_i32(
453 min: i32,
454 max: i32,
455 method: SamplingMethod,
456 )
457 -> Outcome<i32>
458 {
459 if min > max {
460 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
461 }
462
463 Ok(match method {
464 SamplingMethod::Uniform => Self::in_range(min, max),
465 SamplingMethod::GaussianClampedDerived => {
466 let mean = (min as f64 + max as f64) / 2.0;
467 let stdev = (max as f64 - min as f64) / 6.0;
468 let sample = Self::normal_f64(mean, stdev).round();
469 sample.max(min as f64).min(max as f64) as i32
470 }
471 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
472 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
473 sample.max(min as f64).min(max as f64) as i32
474 }
475 })
476 }
477
478 pub fn sample_i64(
479 min: i64,
480 max: i64,
481 method: SamplingMethod,
482 )
483 -> Outcome<i64>
484 {
485 if min > max {
486 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
487 }
488
489 Ok(match method {
490 SamplingMethod::Uniform => Self::in_range(min, max),
491 SamplingMethod::GaussianClampedDerived => {
492 let mean = (min as f64 + max as f64) / 2.0;
493 let stdev = (max as f64 - min as f64) / 6.0;
494 let sample = Self::normal_f64(mean, stdev).round();
495 sample.max(min as f64).min(max as f64) as i64
496 }
497 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
498 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
499 sample.max(min as f64).min(max as f64) as i64
500 }
501 })
502 }
503
504 pub fn sample_i128(
505 min: i128,
506 max: i128,
507 method: SamplingMethod,
508 )
509 -> Outcome<i128>
510 {
511 if min > max {
512 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
513 }
514
515 Ok(match method {
516 SamplingMethod::Uniform => Self::in_range(min, max),
517 SamplingMethod::GaussianClampedDerived => {
518 let min_f64 = min as f64;
519 let max_f64 = max as f64;
520 let mean = (min_f64 + max_f64) / 2.0;
521 let stdev = (max_f64 - min_f64) / 6.0;
522 let sample = Self::normal_f64(mean, stdev).round();
523
524 if sample < i128::MIN as f64 {
525 min
526 } else if sample > i128::MAX as f64 {
527 max
528 } else {
529 sample.max(min_f64).min(max_f64) as i128
530 }
531 }
532 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
533 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
534 if sample < i128::MIN as f64 {
535 min
536 } else if sample > i128::MAX as f64 {
537 max
538 } else {
539 sample.max(min as f64).min(max as f64) as i128
540 }
541 }
542 })
543 }
544
545 pub fn sample_isize(
546 min: isize,
547 max: isize,
548 method: SamplingMethod,
549 )
550 -> Outcome<isize>
551 {
552 if min > max {
553 return Err(err!("Invalid range: {} > {}", min, max; Invalid, Range));
554 }
555
556 Ok(match method {
557 SamplingMethod::Uniform => Self::in_range(min, max),
558 SamplingMethod::GaussianClampedDerived => {
559 let mean = (min as f64 + max as f64) / 2.0;
560 let stdev = (max as f64 - min as f64) / 6.0;
561 let sample = Self::normal_f64(mean, stdev).round();
562 sample.max(min as f64).min(max as f64) as isize
563 }
564 SamplingMethod::GaussianClampedExplicit { mean, stdev } => {
565 let sample = Self::normal_f64(mean as f64, stdev as f64).round();
566 sample.max(min as f64).min(max as f64) as isize
567 }
568 })
569 }
570}
571
572#[cfg(test)]
573mod tests {
574 use super::*;
575
576 #[test]
577 fn test_sample_ranges() -> Outcome<()> {
578 // Test u32.
579 for _ in 0..100 {
580 let val = res!(Rand::sample_u32(10, 100, SamplingMethod::Uniform));
581 req!(val >= 10 && val <= 100, true);
582 }
583
584 // Test i32 negative range.
585 for _ in 0..100 {
586 let val = res!(Rand::sample_i32(-50, 50, SamplingMethod::GaussianClampedDerived));
587 req!(val >= -50 && val <= 50, true);
588 }
589
590 Ok(())
591 }
592
593 #[test]
594 fn test_normal_distribution() -> Outcome<()> {
595 let mean = 100.0;
596 let stdev = 15.0;
597 let mut samples = Vec::new();
598
599 for _ in 0..1000 {
600 samples.push(Rand::normal(mean, stdev));
601 }
602
603 let sample_mean = samples.iter().sum::<f32>() / samples.len() as f32;
604 let variance = samples.iter()
605 .map(|x| (x - sample_mean).powi(2))
606 .sum::<f32>() / (samples.len() - 1) as f32;
607 let sample_stdev = variance.sqrt();
608
609 msg!("Expected mean: {}, Sample mean: {}", mean, sample_mean);
610 msg!("Expected stdev: {}, Sample stdev: {}", stdev, sample_stdev);
611
612 assert!((sample_mean - mean).abs() < 1.0, "Sample mean too far from expected");
613 assert!((sample_stdev - stdev).abs() < 1.0, "Sample stdev too far from expected");
614
615 Ok(())
616 }
617}