1use super::{Error, Weight};
10use crate::Rng;
11use crate::distr::Distribution;
12use crate::distr::uniform::{SampleBorrow, SampleUniform, UniformSampler};
13
14use alloc::vec::Vec;
16use core::fmt::{self, Debug};
17
18#[cfg(feature = "serde")]
19use serde::{Deserialize, Serialize};
20
21#[derive(Debug, Clone, PartialEq)]
80#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
81pub struct WeightedIndex<X: SampleUniform + PartialOrd> {
82 cumulative_weights: Vec<X>,
83 total_weight: X,
84 weight_distribution: X::Sampler,
85}
86
87impl<X: SampleUniform + PartialOrd> WeightedIndex<X> {
88 pub fn new<I>(weights: I) -> Result<WeightedIndex<X>, Error>
100 where
101 I: IntoIterator,
102 I::Item: SampleBorrow<X>,
103 X: Weight,
104 {
105 let mut iter = weights.into_iter();
106 let mut total_weight: X = iter.next().ok_or(Error::InvalidInput)?.borrow().clone();
107
108 let zero = X::ZERO;
109 if !(total_weight >= zero) {
110 return Err(Error::InvalidWeight);
111 }
112
113 let mut weights = Vec::<X>::with_capacity(iter.size_hint().0);
114 for w in iter {
115 if !(w.borrow() >= &zero) {
118 return Err(Error::InvalidWeight);
119 }
120 weights.push(total_weight.clone());
121
122 if let Err(()) = total_weight.checked_add_assign(w.borrow()) {
123 return Err(Error::Overflow);
124 }
125 }
126
127 if total_weight == zero {
128 return Err(Error::InsufficientNonZero);
129 }
130 let distr = X::Sampler::new(zero, total_weight.clone()).map_err(|_| Error::Overflow)?;
131
132 Ok(WeightedIndex {
133 cumulative_weights: weights,
134 total_weight,
135 weight_distribution: distr,
136 })
137 }
138
139 pub fn update_weights(&mut self, new_weights: &[(usize, &X)]) -> Result<(), Error>
161 where
162 X: for<'a> core::ops::AddAssign<&'a X>
163 + for<'a> core::ops::SubAssign<&'a X>
164 + Clone
165 + Default,
166 {
167 if new_weights.is_empty() {
168 return Ok(());
169 }
170
171 let zero = <X as Default>::default();
172
173 let mut total_weight = self.total_weight.clone();
174
175 let mut prev_i = None;
178 for &(i, w) in new_weights {
179 if let Some(old_i) = prev_i {
180 if old_i >= i {
181 return Err(Error::InvalidInput);
182 }
183 }
184 if !(*w >= zero) {
185 return Err(Error::InvalidWeight);
186 }
187 if i > self.cumulative_weights.len() {
188 return Err(Error::InvalidInput);
189 }
190
191 let mut old_w = if i < self.cumulative_weights.len() {
192 self.cumulative_weights[i].clone()
193 } else {
194 self.total_weight.clone()
195 };
196 if i > 0 {
197 old_w -= &self.cumulative_weights[i - 1];
198 }
199
200 total_weight -= &old_w;
201 total_weight += w;
202 prev_i = Some(i);
203 }
204 if total_weight <= zero {
205 return Err(Error::InsufficientNonZero);
206 }
207 let weight_distribution =
208 X::Sampler::new(zero.clone(), total_weight.clone()).map_err(|_| Error::Overflow)?;
209
210 let mut iter = new_weights.iter();
213
214 let mut prev_weight = zero.clone();
215 let mut next_new_weight = iter.next();
216 let &(first_new_index, _) = next_new_weight.unwrap();
217 let mut cumulative_weight = if first_new_index > 0 {
218 self.cumulative_weights[first_new_index - 1].clone()
219 } else {
220 zero.clone()
221 };
222 for i in first_new_index..self.cumulative_weights.len() {
223 match next_new_weight {
224 Some(&(j, w)) if i == j => {
225 cumulative_weight += w;
226 next_new_weight = iter.next();
227 }
228 _ => {
229 let mut tmp = self.cumulative_weights[i].clone();
230 tmp -= &prev_weight; cumulative_weight += &tmp;
232 }
233 }
234 prev_weight = cumulative_weight.clone();
235 core::mem::swap(&mut prev_weight, &mut self.cumulative_weights[i]);
236 }
237
238 self.total_weight = total_weight;
239 self.weight_distribution = weight_distribution;
240
241 Ok(())
242 }
243}
244
245pub struct WeightedIndexIter<'a, X: SampleUniform + PartialOrd> {
248 weighted_index: &'a WeightedIndex<X>,
249 index: usize,
250}
251
252impl<X> Debug for WeightedIndexIter<'_, X>
253where
254 X: SampleUniform + PartialOrd + Debug,
255 X::Sampler: Debug,
256{
257 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
258 f.debug_struct("WeightedIndexIter")
259 .field("weighted_index", &self.weighted_index)
260 .field("index", &self.index)
261 .finish()
262 }
263}
264
265impl<X> Clone for WeightedIndexIter<'_, X>
266where
267 X: SampleUniform + PartialOrd,
268{
269 fn clone(&self) -> Self {
270 WeightedIndexIter {
271 weighted_index: self.weighted_index,
272 index: self.index,
273 }
274 }
275}
276
277impl<X> Iterator for WeightedIndexIter<'_, X>
278where
279 X: for<'b> core::ops::SubAssign<&'b X> + SampleUniform + PartialOrd + Clone,
280{
281 type Item = X;
282
283 fn next(&mut self) -> Option<Self::Item> {
284 match self.weighted_index.weight(self.index) {
285 None => None,
286 Some(weight) => {
287 self.index += 1;
288 Some(weight)
289 }
290 }
291 }
292
293 fn size_hint(&self) -> (usize, Option<usize>) {
294 let remaining = self.weighted_index.cumulative_weights.len() + 1 - self.index;
295 (remaining, Some(remaining))
296 }
297}
298
299impl<X> ExactSizeIterator for WeightedIndexIter<'_, X> where
300 X: for<'b> core::ops::SubAssign<&'b X> + SampleUniform + PartialOrd + Clone
301{
302}
303
304impl<X: SampleUniform + PartialOrd + Clone> WeightedIndex<X> {
305 pub fn weight(&self, index: usize) -> Option<X>
322 where
323 X: for<'a> core::ops::SubAssign<&'a X>,
324 {
325 let mut weight = if let Some(weight) = self.cumulative_weights.get(index) {
326 weight.clone()
327 } else if index == self.cumulative_weights.len() {
328 self.total_weight.clone()
329 } else {
330 return None;
331 };
332
333 if index > 0 {
334 weight -= &self.cumulative_weights[index - 1];
335 }
336 Some(weight)
337 }
338
339 pub fn weights(&self) -> WeightedIndexIter<'_, X>
356 where
357 X: for<'a> core::ops::SubAssign<&'a X>,
358 {
359 WeightedIndexIter {
360 weighted_index: self,
361 index: 0,
362 }
363 }
364
365 pub fn total_weight(&self) -> X {
367 self.total_weight.clone()
368 }
369}
370
371impl<X> Distribution<usize> for WeightedIndex<X>
372where
373 X: SampleUniform + PartialOrd,
374{
375 fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> usize {
376 let chosen_weight = self.weight_distribution.sample(rng);
377 self.cumulative_weights
379 .partition_point(|w| w <= &chosen_weight)
380 }
381}
382
383#[cfg(test)]
384mod test {
385 use super::*;
386 use crate::RngExt;
387
388 #[cfg(feature = "serde")]
389 #[test]
390 fn test_weightedindex_serde() {
391 let weighted_index = WeightedIndex::new([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).unwrap();
392
393 let ser_weighted_index = postcard::to_allocvec(&weighted_index).unwrap();
394 let de_weighted_index: WeightedIndex<i32> =
395 postcard::from_bytes(&ser_weighted_index).unwrap();
396
397 assert_eq!(
398 de_weighted_index.cumulative_weights,
399 weighted_index.cumulative_weights
400 );
401 assert_eq!(de_weighted_index.total_weight, weighted_index.total_weight);
402 }
403
404 #[test]
405 fn test_accepting_nan() {
406 assert_eq!(
407 WeightedIndex::new([f32::NAN, 0.5]).unwrap_err(),
408 Error::InvalidWeight,
409 );
410 assert_eq!(
411 WeightedIndex::new([f32::NAN]).unwrap_err(),
412 Error::InvalidWeight,
413 );
414 assert_eq!(
415 WeightedIndex::new([0.5, f32::NAN]).unwrap_err(),
416 Error::InvalidWeight,
417 );
418
419 assert_eq!(
420 WeightedIndex::new([0.5, 7.0])
421 .unwrap()
422 .update_weights(&[(0, &f32::NAN)])
423 .unwrap_err(),
424 Error::InvalidWeight,
425 )
426 }
427
428 #[test]
429 #[cfg_attr(miri, ignore)] fn test_weightedindex() {
431 let mut r = crate::test::rng(700);
432 const N_REPS: u32 = 5000;
433 let weights = [1u32, 2, 3, 0, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7];
434 let total_weight = weights.iter().sum::<u32>() as f32;
435
436 let verify = |result: [i32; 14]| {
437 for (i, count) in result.iter().enumerate() {
438 let exp = (weights[i] * N_REPS) as f32 / total_weight;
439 let mut err = (*count as f32 - exp).abs();
440 if err != 0.0 {
441 err /= exp;
442 }
443 assert!(err <= 0.25);
444 }
445 };
446
447 let mut chosen = [0i32; 14];
449 let distr = WeightedIndex::new(weights.to_vec()).unwrap();
450 for _ in 0..N_REPS {
451 chosen[distr.sample(&mut r)] += 1;
452 }
453 verify(chosen);
454
455 chosen = [0i32; 14];
457 let distr = WeightedIndex::new(&weights[..]).unwrap();
458 for _ in 0..N_REPS {
459 chosen[distr.sample(&mut r)] += 1;
460 }
461 verify(chosen);
462
463 chosen = [0i32; 14];
465 let distr = WeightedIndex::new(weights.iter()).unwrap();
466 for _ in 0..N_REPS {
467 chosen[distr.sample(&mut r)] += 1;
468 }
469 verify(chosen);
470
471 for _ in 0..5 {
472 assert_eq!(WeightedIndex::new([0, 1]).unwrap().sample(&mut r), 1);
473 assert_eq!(WeightedIndex::new([1, 0]).unwrap().sample(&mut r), 0);
474 assert_eq!(
475 WeightedIndex::new([0, 0, 0, 0, 10, 0])
476 .unwrap()
477 .sample(&mut r),
478 4
479 );
480 }
481 }
482
483 #[test]
484 fn weighted_index_new_errors() {
485 assert_eq!(
486 WeightedIndex::new(&[10][0..0]).unwrap_err(),
487 Error::InvalidInput
488 );
489 assert_eq!(
490 WeightedIndex::new([0]).unwrap_err(),
491 Error::InsufficientNonZero
492 );
493 assert_eq!(
494 WeightedIndex::new([10, 20, -1, 30]).unwrap_err(),
495 Error::InvalidWeight
496 );
497 assert_eq!(
498 WeightedIndex::new([-10, 20, 1, 30]).unwrap_err(),
499 Error::InvalidWeight
500 );
501 assert_eq!(WeightedIndex::new([-10]).unwrap_err(), Error::InvalidWeight);
502 assert_eq!(
503 WeightedIndex::new([f64::INFINITY]).unwrap_err(),
504 Error::Overflow
505 );
506 }
507
508 #[test]
509 fn test_update_weights() {
510 let data = [
511 (
512 &[10u32, 2, 3, 4][..],
513 &[(1, &100), (2, &4)][..], &[10, 100, 4, 4][..],
515 ),
516 (
517 &[1u32, 2, 3, 0, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7][..],
518 &[(2, &1), (5, &1), (13, &100)][..], &[1u32, 2, 1, 0, 5, 1, 7, 1, 2, 3, 4, 5, 6, 100][..],
520 ),
521 ];
522
523 for (weights, update, expected_weights) in data.iter() {
524 let total_weight = weights.iter().sum::<u32>();
525 let mut distr = WeightedIndex::new(weights.to_vec()).unwrap();
526 assert_eq!(distr.total_weight, total_weight);
527
528 distr.update_weights(update).unwrap();
529 let expected_total_weight = expected_weights.iter().sum::<u32>();
530 let expected_distr = WeightedIndex::new(expected_weights.to_vec()).unwrap();
531 assert_eq!(distr.total_weight, expected_total_weight);
532 assert_eq!(distr.total_weight, expected_distr.total_weight);
533 assert_eq!(distr.cumulative_weights, expected_distr.cumulative_weights);
534 }
535 }
536
537 #[test]
538 fn weighted_index_update_errors() {
539 let mut distr = WeightedIndex::new([1.0, 10.0]).unwrap();
540 assert_eq!(
541 distr.update_weights(&[(0, &f32::INFINITY)]).unwrap_err(),
542 Error::Overflow
543 );
544 }
545
546 #[test]
547 fn test_update_weights_errors() {
548 let data = [
549 (
550 &[1i32, 0, 0][..],
551 &[(0, &0)][..],
552 Error::InsufficientNonZero,
553 ),
554 (
555 &[10, 10, 10, 10][..],
556 &[(1, &-11)][..],
557 Error::InvalidWeight, ),
559 (
560 &[1, 2, 3, 4, 5][..],
561 &[(1, &5), (0, &5)][..], Error::InvalidInput,
563 ),
564 (
565 &[1][..],
566 &[(1, &1)][..], Error::InvalidInput,
568 ),
569 ];
570
571 for (weights, update, err) in data.iter() {
572 let total_weight = weights.iter().sum::<i32>();
573 let mut distr = WeightedIndex::new(weights.to_vec()).unwrap();
574 assert_eq!(distr.total_weight, total_weight);
575 match distr.update_weights(update) {
576 Ok(_) => panic!("Expected update_weights to fail, but it succeeded"),
577 Err(e) => assert_eq!(e, *err),
578 }
579 }
580 }
581
582 #[test]
583 fn test_weight_at() {
584 let data = [
585 &[1][..],
586 &[10, 2, 3, 4][..],
587 &[1, 2, 3, 0, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7][..],
588 &[u32::MAX][..],
589 ];
590
591 for weights in data.iter() {
592 let distr = WeightedIndex::new(weights.to_vec()).unwrap();
593 for (i, weight) in weights.iter().enumerate() {
594 assert_eq!(distr.weight(i), Some(*weight));
595 }
596 assert_eq!(distr.weight(weights.len()), None);
597 assert_eq!(distr.weight(usize::MAX), None);
598 }
599 }
600
601 #[test]
602 fn test_weights() {
603 let data = [
604 &[1][..],
605 &[10, 2, 3, 4][..],
606 &[1, 2, 3, 0, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7][..],
607 &[u32::MAX][..],
608 ];
609
610 for weights in data.iter() {
611 let distr = WeightedIndex::new(weights.to_vec()).unwrap();
612 assert_eq!(distr.weights().collect::<Vec<_>>(), weights.to_vec());
613
614 let mut iter = distr.weights();
615 for (index, expected) in weights.iter().enumerate() {
616 let remaining = weights.len() - index;
617 assert_eq!(iter.size_hint(), (remaining, Some(remaining)));
618 assert_eq!(iter.len(), remaining);
619 assert_eq!(iter.clone().collect::<Vec<_>>(), weights[index..]);
620 assert_eq!(iter.next(), Some(*expected));
621 }
622 for _ in 0..2 {
623 assert_eq!(iter.size_hint(), (0, Some(0)));
624 assert_eq!(iter.len(), 0);
625 assert_eq!(iter.next(), None);
626 }
627 }
628 }
629
630 #[test]
631 fn value_stability() {
632 fn test_samples<X: Weight + SampleUniform + PartialOrd, I>(
633 weights: I,
634 buf: &mut [usize],
635 expected: &[usize],
636 ) where
637 I: IntoIterator,
638 I::Item: SampleBorrow<X>,
639 {
640 assert_eq!(buf.len(), expected.len());
641 let distr = WeightedIndex::new(weights).unwrap();
642 let mut rng = crate::test::rng(701);
643 for r in buf.iter_mut() {
644 *r = rng.sample(&distr);
645 }
646 assert_eq!(buf, expected);
647 }
648
649 let mut buf = [0; 10];
650 test_samples(
651 [1i32, 1, 1, 1, 1, 1, 1, 1, 1],
652 &mut buf,
653 &[0, 6, 2, 6, 3, 4, 7, 8, 2, 5],
654 );
655 test_samples(
656 [0.7f32, 0.1, 0.1, 0.1],
657 &mut buf,
658 &[0, 0, 0, 1, 0, 0, 2, 3, 0, 0],
659 );
660 test_samples(
661 [1.0f64, 0.999, 0.998, 0.997],
662 &mut buf,
663 &[2, 2, 1, 3, 2, 1, 3, 3, 2, 1],
664 );
665 }
666
667 #[test]
668 fn weighted_index_distributions_can_be_compared() {
669 assert_eq!(WeightedIndex::new([1, 2]), WeightedIndex::new([1, 2]));
670 }
671
672 #[test]
673 fn overflow() {
674 assert_eq!(WeightedIndex::new([2, usize::MAX]), Err(Error::Overflow));
675 }
676
677 #[test]
678 fn overflow_float() {
679 assert_eq!(
680 WeightedIndex::new([f64::MAX, f64::MAX]),
681 Err(Error::Overflow)
682 );
683 assert_eq!(
684 WeightedIndex::new([f32::MAX, f32::MAX]),
685 Err(Error::Overflow)
686 );
687 assert_eq!(WeightedIndex::new([f64::INFINITY]), Err(Error::Overflow));
688
689 let mut distr = WeightedIndex::new([1.0f64, 2.0]).unwrap();
691 assert_eq!(
692 distr.update_weights(&[(0, &f64::MAX), (1, &f64::MAX)]),
693 Err(Error::Overflow)
694 );
695 assert_eq!(distr, WeightedIndex::new([1.0f64, 2.0]).unwrap());
696 }
697}