diff --git a/CHANGELOG.md b/CHANGELOG.md index 57020070f7..e1d3bd3dcf 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,9 +15,13 @@ You may also find the [Upgrade Guide](https://rust-random.github.io/book/update. - Fix spurious `Error::NonFinite` from `Uniform::new_inclusive` on large finite float ranges such as `0.0..=f64::MAX` ([#1821]) - Fix possible panic due to sampling a deserialized `Uniform` ([#1831]) +### Changes +- Report exact remaining lengths from `WeightedIndex::weights()` and reduce overhead when reading weights ([#1838]) + [#1808]: https://github.com/rust-random/rand/pull/1808 [#1821]: https://github.com/rust-random/rand/pull/1821 [#1831]: https://github.com/rust-random/rand/pull/1831 +[#1838]: https://github.com/rust-random/rand/pull/1838 ## [0.10.2] — 2026-07-02 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 7ca3ee1ad9..6262b49bf0 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 { @@ -584,6 +594,7 @@ mod test { assert_eq!(distr.weight(i), Some(*weight)); } assert_eq!(distr.weight(weights.len()), None); + assert_eq!(distr.weight(usize::MAX), None); } } @@ -599,6 +610,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); + } } }