use rten_tensor::RandomSource;
use rten_tensor::rng::XorShiftRng;
pub struct ReducedRangeRng {
reduce_range: bool,
rng: XorShiftRng,
}
impl ReducedRangeRng {
pub fn new(reduce_range: bool, seed: u64) -> Self {
Self {
rng: XorShiftRng::new(seed),
reduce_range,
}
}
}
impl RandomSource<i8> for ReducedRangeRng {
fn next(&mut self) -> i8 {
if self.reduce_range {
((self.rng.next_u64() % 128) as i16 - 64i16) as i8
} else {
self.rng.next_u64() as i8
}
}
}
impl RandomSource<u8> for ReducedRangeRng {
fn next(&mut self) -> u8 {
if self.reduce_range {
(self.rng.next_u64() % 128) as u8
} else {
self.rng.next_u64() as u8
}
}
}
#[cfg(test)]
mod tests {
use rten_tensor::RandomSource;
use super::ReducedRangeRng;
#[test]
fn test_reduced_range_rng() {
let mut rng = ReducedRangeRng::new(true, 1234);
for _ in 0..100 {
let x: i8 = rng.next();
assert!(x >= -64 && x <= 63);
let x: u8 = rng.next();
assert!(x <= 127);
}
}
}