Skip to main content

rand/distr/weighted/
weighted_index.rs

1// Copyright 2018 Developers of the Rand project.
2//
3// Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
4// https://www.apache.org/licenses/LICENSE-2.0> or the MIT license
5// <LICENSE-MIT or https://opensource.org/licenses/MIT>, at your
6// option. This file may not be copied, modified, or distributed
7// except according to those terms.
8
9use super::{Error, Weight};
10use crate::Rng;
11use crate::distr::Distribution;
12use crate::distr::uniform::{SampleBorrow, SampleUniform, UniformSampler};
13
14// Note that this whole module is only imported if feature="alloc" is enabled.
15use alloc::vec::Vec;
16use core::fmt::{self, Debug};
17
18#[cfg(feature = "serde")]
19use serde::{Deserialize, Serialize};
20
21/// A distribution using weighted sampling of discrete items.
22///
23/// Sampling a `WeightedIndex` distribution returns the index of a randomly
24/// selected element from the iterator used when the `WeightedIndex` was
25/// created. The chance of a given element being picked is proportional to the
26/// weight of the element. The weights can use any type `X` for which an
27/// implementation of [`Uniform<X>`] exists. The implementation guarantees that
28/// elements with zero weight are never picked, even when the weights are
29/// floating point numbers.
30///
31/// # Performance
32///
33/// Time complexity of sampling from `WeightedIndex` is `O(log N)` where
34/// `N` is the number of weights.
35/// See also [`rand_distr::weighted`] for alternative implementations supporting
36/// potentially-faster sampling or a more easily modifiable tree structure.
37///
38/// A `WeightedIndex<X>` contains a `Vec<X>` and a [`Uniform<X>`] and so its
39/// size is the sum of the size of those objects, possibly plus some alignment.
40///
41/// Creating a `WeightedIndex<X>` will allocate enough space to hold `N - 1`
42/// weights of type `X`, where `N` is the number of weights. However, since
43/// `Vec` doesn't guarantee a particular growth strategy, additional memory
44/// might be allocated but not used. Since the `WeightedIndex` object also
45/// contains an instance of `X::Sampler`, this might cause additional allocations,
46/// though for primitive types, [`Uniform<X>`] doesn't allocate any memory.
47///
48/// Sampling from `WeightedIndex` will result in a single call to
49/// `Uniform<X>::sample` (method of the [`Distribution`] trait), which typically
50/// will request a single value from the underlying [`Rng`], though the
51/// exact number depends on the implementation of `Uniform<X>::sample`.
52///
53/// # Example
54///
55/// ```
56/// use rand::prelude::*;
57/// use rand::distr::weighted::WeightedIndex;
58///
59/// let choices = ['a', 'b', 'c'];
60/// let weights = [2,   1,   1];
61/// let dist = WeightedIndex::new(&weights).unwrap();
62/// let mut rng = rand::rng();
63/// for _ in 0..100 {
64///     // 50% chance to print 'a', 25% chance to print 'b', 25% chance to print 'c'
65///     println!("{}", choices[dist.sample(&mut rng)]);
66/// }
67///
68/// let items = [('a', 0.0), ('b', 3.0), ('c', 7.0)];
69/// let dist2 = WeightedIndex::new(items.iter().map(|item| item.1)).unwrap();
70/// for _ in 0..100 {
71///     // 0% chance to print 'a', 30% chance to print 'b', 70% chance to print 'c'
72///     println!("{}", items[dist2.sample(&mut rng)].0);
73/// }
74/// ```
75///
76/// [`Uniform<X>`]: crate::distr::Uniform
77/// [`Rng`]: crate::Rng
78/// [`rand_distr::weighted`]: https://docs.rs/rand_distr/latest/rand_distr/weighted/index.html
79#[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    /// Creates a new a `WeightedIndex` [`Distribution`] using the values
89    /// in `weights`. The weights can use any type `X` for which an
90    /// implementation of [`Uniform<X>`] exists.
91    ///
92    /// Error cases:
93    /// -   [`Error::InvalidInput`] when the iterator `weights` is empty.
94    /// -   [`Error::InvalidWeight`] when a weight is not-a-number or negative.
95    /// -   [`Error::InsufficientNonZero`] when the sum of all weights is zero.
96    /// -   [`Error::Overflow`] when the sum of all weights overflows.
97    ///
98    /// [`Uniform<X>`]: crate::distr::uniform::Uniform
99    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            // Note that `!(w >= x)` is not equivalent to `w < x` for partially
116            // ordered types due to NaNs which are equal to nothing.
117            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    /// Update a subset of weights, without changing the number of weights.
140    ///
141    /// `new_weights` must be sorted by the index.
142    ///
143    /// Using this method instead of `new` might be more efficient if only a small number of
144    /// weights is modified. No allocations are performed, unless the weight type `X` uses
145    /// allocation internally.
146    ///
147    /// In case of error, `self` is not modified. Error cases:
148    /// -   [`Error::InvalidInput`] when `new_weights` are not ordered by
149    ///     index or an index is too large.
150    /// -   [`Error::InvalidWeight`] when a weight is not-a-number or negative.
151    /// -   [`Error::InsufficientNonZero`] when the sum of all weights is zero.
152    ///     Note that due to floating-point loss of precision, this case is not
153    ///     always correctly detected; usage of a fixed-point weight type may be
154    ///     preferred.
155    /// -   [`Error::Overflow`] when the sum of all weights overflows.
156    ///
157    /// Updates take `O(N)` time. If you need to frequently update weights, consider
158    /// [`rand_distr::weighted_tree`](https://docs.rs/rand_distr/*/rand_distr/weighted_tree/index.html)
159    /// as an alternative where an update is `O(log N)`.
160    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        // Check for errors first, so we don't modify `self` in case something
176        // goes wrong.
177        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        // Update the weights. Because we checked all the preconditions in the
211        // previous loop, this should never panic.
212        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; // We know this is positive.
231                    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
245/// A lazy-loading iterator over the weights of a `WeightedIndex` distribution.
246/// This is returned by [`WeightedIndex::weights`].
247pub 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    /// Returns the weight at the given index, if it exists.
306    ///
307    /// If the index is out of bounds, this will return `None`.
308    ///
309    /// # Example
310    ///
311    /// ```
312    /// use rand::distr::weighted::WeightedIndex;
313    ///
314    /// let weights = [0, 1, 2];
315    /// let dist = WeightedIndex::new(&weights).unwrap();
316    /// assert_eq!(dist.weight(0), Some(0));
317    /// assert_eq!(dist.weight(1), Some(1));
318    /// assert_eq!(dist.weight(2), Some(2));
319    /// assert_eq!(dist.weight(3), None);
320    /// ```
321    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    /// Returns a lazy-loading iterator containing the current weights of this distribution.
340    ///
341    /// If this distribution has not been updated since its creation, this will return the
342    /// same weights as were passed to `new`.
343    ///
344    /// # Example
345    ///
346    /// ```
347    /// use rand::distr::weighted::WeightedIndex;
348    ///
349    /// let weights = [1, 2, 3];
350    /// let mut dist = WeightedIndex::new(&weights).unwrap();
351    /// assert_eq!(dist.weights().collect::<Vec<_>>(), vec![1, 2, 3]);
352    /// dist.update_weights(&[(0, &2)]).unwrap();
353    /// assert_eq!(dist.weights().collect::<Vec<_>>(), vec![2, 2, 3]);
354    /// ```
355    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    /// Returns the sum of all weights in this distribution.
366    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        // Find the first item which has a weight *higher* than the chosen weight.
378        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)] // Miri is too slow
430    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        // WeightedIndex from vec
448        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        // WeightedIndex from slice
456        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        // WeightedIndex from iterator
464        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)][..], // positive change
514                &[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)][..], // negative change and last element
519                &[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, // A weight is negative
558            ),
559            (
560                &[1, 2, 3, 4, 5][..],
561                &[(1, &5), (0, &5)][..], // Wrong order
562                Error::InvalidInput,
563            ),
564            (
565                &[1][..],
566                &[(1, &1)][..], // Index too large
567                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        // In case of error, self is not modified.
690        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}