Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
627 changes: 461 additions & 166 deletions Cargo.lock

Large diffs are not rendered by default.

6 changes: 3 additions & 3 deletions Cargo.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[package]
name = "weighted_rand"
version = "0.4.2"
version = "0.5.0"
authors = ["Ichi <ichi.h3@gmail.com>"]
description = "A weighted random sampling crate using Walker's Alias Method."
documentation = "https://docs.rs/weighted_rand"
Expand All @@ -15,10 +15,10 @@ edition = "2018"

[dependencies]
serde = { version = "1.0", features = ["derive"] }
rand = "0.8"
rand = " =0.10"

[dev-dependencies]
criterion = "0.3"
criterion = "0.5.1"

[[bench]]
name = "benchmark"
Expand Down
14 changes: 7 additions & 7 deletions benches/benchmark.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
use rand::prelude::*;
use rand::{rngs::ThreadRng, RngExt};
use weighted_rand::builder::*;

use criterion::Criterion;
Expand Down Expand Up @@ -31,7 +31,7 @@ fn bench_generate_by_wam_next_rng(c: &mut Criterion) {
let builder = WalkerTableBuilder::new(&WEIGHTS);
let table = builder.build();

let mut rng = rand::thread_rng();
let mut rng = rand::rng();

let mut result = [0; 100_000];

Expand All @@ -50,8 +50,8 @@ fn bench_generate_by_csm(c: &mut Criterion) {
.collect::<Vec<f32>>()
.to_vec();

let csm = CSM { probs: &probs };
let mut rng = rand::thread_rng();
let csm = CumulativeSumMethod { probs: &probs };
let mut rng = rand::rng();
let mut result = [0; 100_000];

c.bench_function("generate_by_csm", |b| {
Expand All @@ -76,13 +76,13 @@ criterion_main!(benches);

// Weighted random sampling using Cumulative Sum Method

struct CSM<'a> {
struct CumulativeSumMethod<'a> {
probs: &'a [f32],
}

impl CSM<'_> {
impl CumulativeSumMethod<'_> {
fn next(&self, rng: &mut ThreadRng) -> usize {
let r = rng.gen::<f32>();
let r = rng.random::<f32>();
for (i, p) in self.probs.iter().enumerate() {
if r <= *p {
return i;
Expand Down
3 changes: 1 addition & 2 deletions examples/cheating_coin.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
use rand;
use weighted_rand::builder::*;

fn main() {
Expand All @@ -13,7 +12,7 @@ fn main() {
// loops, we recommend using the next_rng method with an
// external ThreadRng instance.
let mut result = [""; 10000];
let mut rng = rand::thread_rng();
let mut rng = rand::rng();
for r in &mut result {
let j = wa_table.next_rng(&mut rng);
*r = cheating_coin[j];
Expand Down
43 changes: 20 additions & 23 deletions src/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
use crate::table::WalkerTable;
use crate::util::math::gcd_for_slice;

#[allow(clippy::new_ret_no_self)]
pub trait NewBuilder<T> {
/// Creates a new instance of [`WalkerTableBuilder`] from
/// [`&[u32]`] or [`&[f32]`].
Expand All @@ -18,11 +19,9 @@ pub trait NewBuilder<T> {
/// ```rust
/// use weighted_rand::builder::*;
///
/// fn main() {
/// let index_weights = [1, 2, 3, 4];
/// let builder = WalkerTableBuilder::new(&index_weights);
/// let wa_table = builder.build();
/// }
/// let index_weights = [1, 2, 3, 4];
/// let builder = WalkerTableBuilder::new(&index_weights);
/// let wa_table = builder.build();
/// ```
///
/// Also, `index_weiaghts` supports [`&[f32]`], like `[0.1, 0.2, 0.3, 0.4]`
Expand Down Expand Up @@ -88,7 +87,7 @@ impl WalkerTableBuilder {

/// Calculates the sum of `index_weights`.
fn sum(&self) -> u32 {
self.index_weights.iter().fold(0, |acc, cur| acc + cur)
self.index_weights.iter().sum::<u32>()
}

/// Calculates the mean of `index_weights`.
Expand All @@ -104,24 +103,19 @@ impl WalkerTableBuilder {

let mut aliases = vec![0; table_len];
let mut probs = vec![0.0; table_len];
loop {
match below_vec.pop() {
Some(below) => {
if let Some(above) = above_vec.pop() {
let diff = mean - below.1;
aliases[below.0] = above.0 as usize;
probs[below.0] = diff as f32 / mean as f32;
if above.1 - diff <= mean {
below_vec.push((above.0, above.1 - diff));
} else {
above_vec.push((above.0, above.1 - diff));
}
} else {
aliases[below.0] = below.0 as usize;
probs[below.0] = below.1 as f32 / mean as f32;
}
while let Some(below) = below_vec.pop() {
if let Some(above) = above_vec.pop() {
let diff = mean - below.1;
aliases[below.0] = above.0;
probs[below.0] = diff as f32 / mean as f32;
if above.1 - diff <= mean {
below_vec.push((above.0, above.1 - diff));
} else {
above_vec.push((above.0, above.1 - diff));
}
None => break,
} else {
aliases[below.0] = below.0;
probs[below.0] = below.1 as f32 / mean as f32;
}
}

Expand All @@ -131,6 +125,7 @@ impl WalkerTableBuilder {
/// Divide the values of `index_weights` based on the mean of them.
///
/// The tail value is a weight and head is its index.
#[allow(clippy::type_complexity)]
fn separate_weight(&self) -> (Vec<(usize, u32)>, Vec<(usize, u32)>) {
let mut below_vec = Vec::with_capacity(self.index_weights.len());
let mut above_vec = Vec::with_capacity(self.index_weights.len());
Expand All @@ -151,6 +146,7 @@ mod builder_test {
use crate::table::WalkerTable;

#[test]
#[allow(clippy::excessive_precision)]
fn make_table_from_u32() {
let index_weights = [2, 7, 9, 2, 4, 8, 1, 3, 6, 5];
let builder = WalkerTableBuilder::new(&index_weights);
Expand All @@ -176,6 +172,7 @@ mod builder_test {
}

#[test]
#[allow(clippy::excessive_precision)]
fn make_table_from_f32() {
let index_weights = [0.1, 0.2, 0.3, -0.4];
let builder = WalkerTableBuilder::new(&index_weights);
Expand Down
66 changes: 31 additions & 35 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,22 +21,20 @@
//! ```rust
//! use weighted_rand::builder::*;
//!
//! fn main() {
//! let fruit = ["Apple", "Banana", "Orange", "Peach"];
//!
//! // Define the weights for each index corresponding
//! // to the above list.
//! // In the following case, the ratio of each weight
//! // is "2 : 1 : 7 : 0", and the output probabilities
//! // for each index are 0.2, 0.1, 0.7 and 0.
//! let index_weights = [2, 1, 7, 0];
//!
//! let builder = WalkerTableBuilder::new(&index_weights);
//! let wa_table = builder.build();
//!
//! for i in (0..10).map(|_| wa_table.next()) {
//! println!("{}", fruit[i]);
//! }
//! let fruit = ["Apple", "Banana", "Orange", "Peach"];
//!
//! // Define the weights for each index corresponding
//! // to the above list.
//! // In the following case, the ratio of each weight
//! // is "2 : 1 : 7 : 0", and the output probabilities
//! // for each index are 0.2, 0.1, 0.7 and 0.
//! let index_weights = [2, 1, 7, 0];
//!
//! let builder = WalkerTableBuilder::new(&index_weights);
//! let wa_table = builder.build();
//!
//! for i in (0..10).map(|_| wa_table.next()) {
//! println!("{}", fruit[i]);
//! }
//! ```
//!
Expand All @@ -46,26 +44,24 @@
//! use rand;
//! use weighted_rand::builder::*;
//!
//! fn main() {
//! // Coin with a 5% higher probability of heads than tails
//! let cheating_coin = ["Heads!", "Tails!"];
//! let index_weights = [0.55, 0.45];
//!
//! let builder = WalkerTableBuilder::new(&index_weights);
//! let wa_table = builder.build();
//!
//! // If you want to process something in a large number of
//! // loops, we recommend using the next_rng method with an
//! // external ThreadRng instance.
//! let mut result = [""; 10000];
//! let mut rng = rand::thread_rng();
//! for r in &mut result {
//! let j = wa_table.next_rng(&mut rng);
//! *r = cheating_coin[j];
//! }
//!
//! // println!("{:?}", result);
//! // Coin with a 5% higher probability of heads than tails
//! let cheating_coin = ["Heads!", "Tails!"];
//! let index_weights = [0.55, 0.45];
//!
//! let builder = WalkerTableBuilder::new(&index_weights);
//! let wa_table = builder.build();
//!
//! // If you want to process something in a large number of
//! // loops, we recommend using the next_rng method with an
//! // external ThreadRng instance.
//! let mut result = [""; 10000];
//! let mut rng = rand::thread_rng();
//! for r in &mut result {
//! let j = wa_table.next_rng(&mut rng);
//! *r = cheating_coin[j];
//! }
//!
//! // println!("{:?}", result);
//! ```
//!

Expand Down
13 changes: 5 additions & 8 deletions src/table.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,22 +28,19 @@ pub struct WalkerTable {
impl WalkerTable {
/// Creates a new instance of [`WalkerTable`].
pub fn new(aliases: Vec<usize>, probs: Vec<f32>) -> WalkerTable {
WalkerTable {
aliases: aliases,
probs: probs,
}
WalkerTable { aliases, probs }
}

/// Returns an index at random.
pub fn next(&self) -> usize {
let mut rng = rand::thread_rng();
let mut rng = rand::rng();
self.next_rng(&mut rng)
}

/// Returns an index at random using an external RNG which implements Rng.
pub fn next_rng(&self, rng: &mut impl Rng) -> usize {
let i = rng.gen_range(0..self.probs.len());
let r = rng.gen::<f32>();
let i = rng.random_range(0..self.probs.len());
let r = rng.random::<f32>();
if r < self.probs[i] {
return self.aliases[i];
}
Expand All @@ -70,7 +67,7 @@ mod table_test {
let builder = WalkerTableBuilder::new(&index_weights);
let wa_table = builder.build();

let mut rng = rand::thread_rng();
let mut rng = rand::rng();

let idxs = (0..N)
.map(|_| wa_table.next_rng(&mut rng))
Expand Down
10 changes: 5 additions & 5 deletions src/util.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ pub mod math {
fn gcd(a: u32, b: u32) -> u32 {
let (a, b) = if a < b { (b, a) } else { (a, b) };

if a % b == 0 {
if a.is_multiple_of(b) {
b
} else {
gcd(b, a % b)
Expand All @@ -17,10 +17,11 @@ pub mod math {
let mut iter = slice.iter().skip_while(|x| x == &&0);
let first = match iter.next() {
Some(v) => *v,
None => return 1
None => return 1,
};

let gcd = iter.fold(
// create gcd
iter.fold(
first,
|acc, cur| {
if *cur == 0 {
Expand All @@ -29,8 +30,7 @@ pub mod math {
gcd(*cur, acc)
}
},
);
gcd
)
}
}

Expand Down