From 777e43f279ed42e616792b2e72b3d0a825b299c9 Mon Sep 17 00:00:00 2001 From: Vishakha Ojha Date: Sat, 30 May 2026 18:47:23 +0000 Subject: [PATCH 1/2] Implement uniform noaml and integers distributions Signed-off-by: Vishakha Ojha --- crates/mohu-random/src/continuous.rs | 52 ++++++++++++++++++++++++++++ crates/mohu-random/src/discrete.rs | 24 +++++++++++++ 2 files changed, 76 insertions(+) diff --git a/crates/mohu-random/src/continuous.rs b/crates/mohu-random/src/continuous.rs index dc03763..b61b31b 100644 --- a/crates/mohu-random/src/continuous.rs +++ b/crates/mohu-random/src/continuous.rs @@ -1 +1,53 @@ // continuous — implementation pending + +use mohu_buffer::Buffer; +use mohu_error::{MohuError, MohuResult}; +use rand::Rng; + +pub fn uniform(shape: &[usize], low: f64, high: f64) -> MohuResult { + if high <= low { + return Err(MohuError::domain( + "uniform", + "high must be greater than low", + )); + } + + let n: usize = shape.iter().product(); + + let mut rng = rand::rng(); + + let data: Vec = (0..n) + .map(|_| rng.random_range(low..high)) + .collect(); + + let buf = Buffer::from_vec(data)?; + buf.reshape(shape) +} + +pub fn normal(shape: &[usize], mean: f64, std: f64) -> MohuResult { + if std <= 0.0 { + return Err(MohuError::domain( + "normal", + "std must be positive", + )); + } + + let n: usize = shape.iter().product(); + let mut rng = rand::rng(); + + let data: Vec = (0..n) + .map(|_| { + let u1: f64 = rng.random(); + let u2: f64 = rng.random(); + + let z0 = + (-2.0 * u1.ln()).sqrt() + * (2.0 * std::f64::consts::PI * u2).cos(); + + mean + std * z0 + }) + .collect(); + + let buf = Buffer::from_vec(data)?; + buf.reshape(shape) +} \ No newline at end of file diff --git a/crates/mohu-random/src/discrete.rs b/crates/mohu-random/src/discrete.rs index 17f166e..76d6b72 100644 --- a/crates/mohu-random/src/discrete.rs +++ b/crates/mohu-random/src/discrete.rs @@ -1 +1,25 @@ // discrete — implementation pending + +use mohu_buffer::Buffer; +use mohu_error::{MohuError, MohuResult}; +use rand::Rng; + +pub fn integers(shape: &[usize], low: i64, high: i64) -> MohuResult { + if high <= low { + return Err(MohuError::domain( + "integers", + "high must be greater than low", + )); + } + + let n: usize = shape.iter().product(); + + let mut rng = rand::rng(); + + let data: Vec = (0..n) + .map(|_| rng.random_range(low..high)) + .collect(); + + let buf = Buffer::from_vec(data)?; + buf.reshape(shape) +} \ No newline at end of file From 36b223f1b508c3cff21dbd30620758f39a835eb0 Mon Sep 17 00:00:00 2001 From: Vishakha Ojha Date: Sat, 6 Jun 2026 18:47:09 +0000 Subject: [PATCH 2/2] fix: use MohuError::domain and ensure! macro, guard box-muller u1>0 Signed-off-by: Vishakha Ojha --- crates/mohu-random/src/continuous.rs | 50 +++++++++++++--------------- crates/mohu-random/src/discrete.rs | 24 +++++-------- 2 files changed, 31 insertions(+), 43 deletions(-) diff --git a/crates/mohu-random/src/continuous.rs b/crates/mohu-random/src/continuous.rs index b61b31b..3205af1 100644 --- a/crates/mohu-random/src/continuous.rs +++ b/crates/mohu-random/src/continuous.rs @@ -1,53 +1,49 @@ -// continuous — implementation pending - use mohu_buffer::Buffer; -use mohu_error::{MohuError, MohuResult}; +use mohu_error::{MohuError, MohuResult, ensure}; use rand::Rng; pub fn uniform(shape: &[usize], low: f64, high: f64) -> MohuResult { - if high <= low { - return Err(MohuError::domain( + ensure!( + low.is_finite() && high.is_finite() && high > low, + MohuError::domain( "uniform", - "high must be greater than low", - )); - } + "high must be greater than low and bounds must be finite" + ) + ); let n: usize = shape.iter().product(); - let mut rng = rand::rng(); - let data: Vec = (0..n) - .map(|_| rng.random_range(low..high)) - .collect(); + let data: Vec = (0..n).map(|_| rng.random_range(low..high)).collect(); - let buf = Buffer::from_vec(data)?; - buf.reshape(shape) + Buffer::from_vec(data)?.reshape(shape) } pub fn normal(shape: &[usize], mean: f64, std: f64) -> MohuResult { - if std <= 0.0 { - return Err(MohuError::domain( - "normal", - "std must be positive", - )); - } + ensure!( + std.is_finite() && std > 0.0, + MohuError::domain("normal", "std must be positive and finite") + ); let n: usize = shape.iter().product(); let mut rng = rand::rng(); let data: Vec = (0..n) .map(|_| { - let u1: f64 = rng.random(); + // Draw u1 from (0, 1] to avoid log(0) = -inf in Box-Muller + let u1: f64 = loop { + let x: f64 = rng.random(); + if x > 0.0 { + break x; + } + }; let u2: f64 = rng.random(); - let z0 = - (-2.0 * u1.ln()).sqrt() - * (2.0 * std::f64::consts::PI * u2).cos(); + let z0 = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos(); mean + std * z0 }) .collect(); - let buf = Buffer::from_vec(data)?; - buf.reshape(shape) -} \ No newline at end of file + Buffer::from_vec(data)?.reshape(shape) +} diff --git a/crates/mohu-random/src/discrete.rs b/crates/mohu-random/src/discrete.rs index 76d6b72..c04418c 100644 --- a/crates/mohu-random/src/discrete.rs +++ b/crates/mohu-random/src/discrete.rs @@ -1,25 +1,17 @@ -// discrete — implementation pending - use mohu_buffer::Buffer; -use mohu_error::{MohuError, MohuResult}; +use mohu_error::{MohuError, MohuResult, ensure}; use rand::Rng; pub fn integers(shape: &[usize], low: i64, high: i64) -> MohuResult { - if high <= low { - return Err(MohuError::domain( - "integers", - "high must be greater than low", - )); - } + ensure!( + high > low, + MohuError::domain("integers", "high must be greater than low") + ); let n: usize = shape.iter().product(); - let mut rng = rand::rng(); - let data: Vec = (0..n) - .map(|_| rng.random_range(low..high)) - .collect(); + let data: Vec = (0..n).map(|_| rng.random_range(low..high)).collect(); - let buf = Buffer::from_vec(data)?; - buf.reshape(shape) -} \ No newline at end of file + Buffer::from_vec(data)?.reshape(shape) +}