From 96aba1169b11537cf0a79202f180cf3a2b2fbbb0 Mon Sep 17 00:00:00 2001 From: Patrik Nilsson <113925545+pnn64@users.noreply.github.com> Date: Mon, 7 Sep 2026 14:19:41 +0200 Subject: [PATCH 1/4] Optimize WeightedIndex weight lookup and iteration Simplify checked weight lookup and implement ExactSizeIterator for weight iteration, avoiding growth reallocations when collecting weights. Add iterator coverage and benchmarks for collection, lookup, and reuse. --- CHANGELOG.md | 3 ++ benches/benches/weighted.rs | 61 +++++++++++++++++++++++++++- src/distr/weighted/weighted_index.rs | 37 ++++++++++++++--- 3 files changed, 93 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 3085994c4d..5a6e48603b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,9 @@ You may also find the [Upgrade Guide](https://rust-random.github.io/book/update. ## [Unreleased] +### Changes +- Report exact remaining lengths from `WeightedIndex::weights()` and reduce overhead when reading weights + ### Fixes - Fix `WeightedIndex` panic when the sum of float weights is infinite; return `Error::Overflow` instead ([#1808]) diff --git a/benches/benches/weighted.rs b/benches/benches/weighted.rs index e5ba371224..a2a05fe942 100644 --- a/benches/benches/weighted.rs +++ b/benches/benches/weighted.rs @@ -6,8 +6,9 @@ // option. This file may not be copied, modified, or distributed // except according to those terms. -use criterion::{Criterion, black_box, criterion_group, criterion_main}; -use rand::distr::weighted::WeightedIndex; +use criterion::{BenchmarkId, Criterion, black_box, criterion_group, criterion_main}; +use rand::distr::uniform::SampleUniform; +use rand::distr::weighted::{Weight, WeightedIndex}; use rand::prelude::*; use rand::seq::index::sample_weighted; @@ -19,6 +20,9 @@ criterion_group!( criterion_main!(benches); pub fn bench(c: &mut Criterion) { + bench_weight_iteration::(c, "u32"); + bench_weight_iteration::(c, "f64"); + c.bench_function("weighted_index_creation", |b| { let mut rng = rand::rng(); let weights = black_box([1u32, 2, 4, 0, 5, 1, 7, 1, 2, 3, 4, 5, 6, 7]); @@ -58,3 +62,56 @@ pub fn bench(c: &mut Criterion) { }); } } + +fn bench_weight_iteration(c: &mut Criterion, name: &str) +where + X: SampleUniform + Weight + PartialOrd + From + core::iter::Sum + for<'a> core::ops::SubAssign<&'a X>, +{ + let mut group = c.benchmark_group(format!("weighted_iter/{name}")); + for length in [1usize, 4, 16, 64, 256, 1024, 16384] { + let distr = WeightedIndex::new((0..length).map(|i| X::from((1 + i % 10) as u32))).unwrap(); + group.bench_function(BenchmarkId::new("collect", length), |b| { + b.iter(|| black_box(&distr).weights().collect::>()) + }); + + // Control cases: neither summing nor reusing capacity needs a size hint. + if [4, 1024].contains(&length) { + group + .bench_function(BenchmarkId::new("sum", length), |b| b.iter(|| black_box(&distr).weights().sum::())); + let mut buffer = Vec::with_capacity(length); + group.bench_function(BenchmarkId::new("reuse", length), |b| { + b.iter(|| { + buffer.clear(); + buffer.extend(black_box(&distr).weights()); + black_box(buffer.as_slice()); + }) + }); + } + + if length == 1024 { + for (position, index) in [ + ("first", 0), + ("middle", length / 2), + ("last", length - 1), + ("past_end", length), + ("max_index", usize::MAX), + ] { + group.bench_function(BenchmarkId::new("weight", position), |b| { + b.iter(|| black_box(&distr).weight(black_box(index))) + }); + } + let mut iter = distr.weights(); + let _ = iter.nth(length / 2 - 1); + group.bench_function(BenchmarkId::new("collect_remaining", length / 2), |b| { + b.iter(|| black_box(iter.clone()).collect::>()) + }); + } + } + let distr = WeightedIndex::new([X::from(1)]).unwrap(); + let mut exhausted = distr.weights(); + let _ = exhausted.next(); + group.bench_function(BenchmarkId::new("collect", 0), |b| { + b.iter(|| black_box(exhausted.clone()).collect::>()) + }); + group.finish(); +} diff --git a/src/distr/weighted/weighted_index.rs b/src/distr/weighted/weighted_index.rs index 0fbc2be361..01071bf60e 100644 --- a/src/distr/weighted/weighted_index.rs +++ b/src/distr/weighted/weighted_index.rs @@ -289,6 +289,16 @@ where } } } + + fn size_hint(&self) -> (usize, Option) { + let remaining = self.weighted_index.cumulative_weights.len() + 1 - self.index; + (remaining, Some(remaining)) + } +} + +impl ExactSizeIterator for WeightedIndexIter<'_, X> where + X: for<'b> core::ops::SubAssign<&'b X> + SampleUniform + PartialOrd + Clone +{ } impl WeightedIndex { @@ -312,12 +322,12 @@ impl WeightedIndex { where X: for<'a> core::ops::SubAssign<&'a X>, { - use core::cmp::Ordering::*; - - let mut weight = match index.cmp(&self.cumulative_weights.len()) { - Less => self.cumulative_weights[index].clone(), - Equal => self.total_weight.clone(), - Greater => return None, + let mut weight = if let Some(weight) = self.cumulative_weights.get(index) { + weight.clone() + } else if index == self.cumulative_weights.len() { + self.total_weight.clone() + } else { + return None; }; if index > 0 { @@ -568,6 +578,7 @@ mod test { assert_eq!(distr.weight(i), Some(*weight)); } assert_eq!(distr.weight(weights.len()), None); + assert_eq!(distr.weight(usize::MAX), None); } } @@ -583,6 +594,20 @@ mod test { for weights in data.iter() { let distr = WeightedIndex::new(weights.to_vec()).unwrap(); assert_eq!(distr.weights().collect::>(), weights.to_vec()); + + let mut iter = distr.weights(); + for (index, expected) in weights.iter().enumerate() { + let remaining = weights.len() - index; + assert_eq!(iter.size_hint(), (remaining, Some(remaining))); + assert_eq!(iter.len(), remaining); + assert_eq!(iter.clone().collect::>(), weights[index..]); + assert_eq!(iter.next(), Some(*expected)); + } + for _ in 0..2 { + assert_eq!(iter.size_hint(), (0, Some(0))); + assert_eq!(iter.len(), 0); + assert_eq!(iter.next(), None); + } } } From d9fcf774962c143ef786f9e933f3d4ab31c65e8d Mon Sep 17 00:00:00 2001 From: Patrik Nilsson <113925545+pnn64@users.noreply.github.com> Date: Sun, 13 Sep 2026 16:46:58 +0200 Subject: [PATCH 2/4] Update CHANGELOG.md --- CHANGELOG.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5a6e48603b..5480239cd9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,12 +11,13 @@ You may also find the [Upgrade Guide](https://rust-random.github.io/book/update. ## [Unreleased] ### Changes -- Report exact remaining lengths from `WeightedIndex::weights()` and reduce overhead when reading weights +- Report exact remaining lengths from `WeightedIndex::weights()` and reduce overhead when reading weights ([#1838]) ### Fixes - Fix `WeightedIndex` panic when the sum of float weights is infinite; return `Error::Overflow` instead ([#1808]) [#1808]: https://github.com/rust-random/rand/pull/1808 +[#1838]: https://github.com/rust-random/rand/pull/1838 ## [0.10.2] — 2026-07-02 From 311051bf74597f6df73d60dcd8bcc3d8a5f6be93 Mon Sep 17 00:00:00 2001 From: Patrik Nilsson <113925545+pnn64@users.noreply.github.com> Date: Sun, 13 Sep 2026 21:56:56 +0200 Subject: [PATCH 3/4] Clarify eager iterator advancement before benchmark timing --- benches/benches/weighted.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/benches/benches/weighted.rs b/benches/benches/weighted.rs index a2a05fe942..cf32f57a08 100644 --- a/benches/benches/weighted.rs +++ b/benches/benches/weighted.rs @@ -101,6 +101,8 @@ where }); } let mut iter = distr.weights(); + // Consume half the weights before timing (nth also consumes its returned item). + // A skip adapter would defer that work until collection. let _ = iter.nth(length / 2 - 1); group.bench_function(BenchmarkId::new("collect_remaining", length / 2), |b| { b.iter(|| black_box(iter.clone()).collect::>()) From bab8a5d24c6f5530825acfd5528b40f84e36a3a7 Mon Sep 17 00:00:00 2001 From: Patrik Nilsson <113925545+pnn64@users.noreply.github.com> Date: Mon, 14 Sep 2026 11:59:11 +0200 Subject: [PATCH 4/4] Remove unnecessary iterator advancement comment --- benches/benches/weighted.rs | 2 -- 1 file changed, 2 deletions(-) diff --git a/benches/benches/weighted.rs b/benches/benches/weighted.rs index cf32f57a08..a2a05fe942 100644 --- a/benches/benches/weighted.rs +++ b/benches/benches/weighted.rs @@ -101,8 +101,6 @@ where }); } let mut iter = distr.weights(); - // Consume half the weights before timing (nth also consumes its returned item). - // A skip adapter would defer that work until collection. let _ = iter.nth(length / 2 - 1); group.bench_function(BenchmarkId::new("collect_remaining", length / 2), |b| { b.iter(|| black_box(iter.clone()).collect::>())