1use alloc::vec::{self, Vec};
11use core::slice;
12use core::{hash::Hash, ops::AddAssign};
13#[cfg(feature = "std")]
15use super::WeightError;
16use crate::distr::uniform::SampleUniform;
17use crate::distr::{Distribution, Uniform};
18use crate::{Rng, RngExt};
19#[cfg(not(feature = "std"))]
20use alloc::collections::BTreeSet;
21#[cfg(feature = "serde")]
22use serde::{Deserialize, Serialize};
23#[cfg(feature = "std")]
24use std::collections::HashSet;
25
26#[cfg(not(any(target_pointer_width = "32", target_pointer_width = "64")))]
27compile_error!("unsupported pointer width");
28
29#[derive(Clone, Debug)]
33#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
34pub enum IndexVec {
35 #[doc(hidden)]
36 U32(Vec<u32>),
37 #[cfg(target_pointer_width = "64")]
38 #[doc(hidden)]
39 U64(Vec<u64>),
40}
41
42impl IndexVec {
43 #[inline]
45 pub fn len(&self) -> usize {
46 match self {
47 IndexVec::U32(v) => v.len(),
48 #[cfg(target_pointer_width = "64")]
49 IndexVec::U64(v) => v.len(),
50 }
51 }
52
53 #[inline]
55 pub fn is_empty(&self) -> bool {
56 match self {
57 IndexVec::U32(v) => v.is_empty(),
58 #[cfg(target_pointer_width = "64")]
59 IndexVec::U64(v) => v.is_empty(),
60 }
61 }
62
63 #[inline]
68 pub fn index(&self, index: usize) -> usize {
69 match self {
70 IndexVec::U32(v) => v[index] as usize,
71 #[cfg(target_pointer_width = "64")]
72 IndexVec::U64(v) => v[index] as usize,
73 }
74 }
75
76 #[inline]
78 pub fn into_vec(self) -> Vec<usize> {
79 match self {
80 IndexVec::U32(v) => v.into_iter().map(|i| i as usize).collect(),
81 #[cfg(target_pointer_width = "64")]
82 IndexVec::U64(v) => v.into_iter().map(|i| i as usize).collect(),
83 }
84 }
85
86 #[inline]
88 pub fn iter(&self) -> IndexVecIter<'_> {
89 match self {
90 IndexVec::U32(v) => IndexVecIter::U32(v.iter()),
91 #[cfg(target_pointer_width = "64")]
92 IndexVec::U64(v) => IndexVecIter::U64(v.iter()),
93 }
94 }
95}
96
97impl IntoIterator for IndexVec {
98 type IntoIter = IndexVecIntoIter;
99 type Item = usize;
100
101 #[inline]
103 fn into_iter(self) -> IndexVecIntoIter {
104 match self {
105 IndexVec::U32(v) => IndexVecIntoIter::U32(v.into_iter()),
106 #[cfg(target_pointer_width = "64")]
107 IndexVec::U64(v) => IndexVecIntoIter::U64(v.into_iter()),
108 }
109 }
110}
111
112impl PartialEq for IndexVec {
113 fn eq(&self, other: &IndexVec) -> bool {
114 use self::IndexVec::*;
115 match (self, other) {
116 (U32(v1), U32(v2)) => v1 == v2,
117 #[cfg(target_pointer_width = "64")]
118 (U64(v1), U64(v2)) => v1 == v2,
119 #[cfg(target_pointer_width = "64")]
120 (U32(v1), U64(v2)) => {
121 (v1.len() == v2.len()) && (v1.iter().zip(v2.iter()).all(|(x, y)| *x as u64 == *y))
122 }
123 #[cfg(target_pointer_width = "64")]
124 (U64(v1), U32(v2)) => {
125 (v1.len() == v2.len()) && (v1.iter().zip(v2.iter()).all(|(x, y)| *x == *y as u64))
126 }
127 }
128 }
129}
130
131impl From<Vec<u32>> for IndexVec {
132 #[inline]
133 fn from(v: Vec<u32>) -> Self {
134 IndexVec::U32(v)
135 }
136}
137
138#[cfg(target_pointer_width = "64")]
139impl From<Vec<u64>> for IndexVec {
140 #[inline]
141 fn from(v: Vec<u64>) -> Self {
142 IndexVec::U64(v)
143 }
144}
145
146#[derive(Debug)]
148pub enum IndexVecIter<'a> {
149 #[doc(hidden)]
150 U32(slice::Iter<'a, u32>),
151 #[cfg(target_pointer_width = "64")]
152 #[doc(hidden)]
153 U64(slice::Iter<'a, u64>),
154}
155
156impl Iterator for IndexVecIter<'_> {
157 type Item = usize;
158
159 #[inline]
160 fn next(&mut self) -> Option<usize> {
161 use self::IndexVecIter::*;
162 match self {
163 U32(iter) => iter.next().map(|i| *i as usize),
164 #[cfg(target_pointer_width = "64")]
165 U64(iter) => iter.next().map(|i| *i as usize),
166 }
167 }
168
169 #[inline]
170 fn size_hint(&self) -> (usize, Option<usize>) {
171 match self {
172 IndexVecIter::U32(v) => v.size_hint(),
173 #[cfg(target_pointer_width = "64")]
174 IndexVecIter::U64(v) => v.size_hint(),
175 }
176 }
177}
178
179impl ExactSizeIterator for IndexVecIter<'_> {}
180
181#[derive(Clone, Debug)]
183pub enum IndexVecIntoIter {
184 #[doc(hidden)]
185 U32(vec::IntoIter<u32>),
186 #[cfg(target_pointer_width = "64")]
187 #[doc(hidden)]
188 U64(vec::IntoIter<u64>),
189}
190
191impl Iterator for IndexVecIntoIter {
192 type Item = usize;
193
194 #[inline]
195 fn next(&mut self) -> Option<Self::Item> {
196 use self::IndexVecIntoIter::*;
197 match self {
198 U32(v) => v.next().map(|i| i as usize),
199 #[cfg(target_pointer_width = "64")]
200 U64(v) => v.next().map(|i| i as usize),
201 }
202 }
203
204 #[inline]
205 fn size_hint(&self) -> (usize, Option<usize>) {
206 use self::IndexVecIntoIter::*;
207 match self {
208 U32(v) => v.size_hint(),
209 #[cfg(target_pointer_width = "64")]
210 U64(v) => v.size_hint(),
211 }
212 }
213}
214
215impl ExactSizeIterator for IndexVecIntoIter {}
216
217#[track_caller]
240pub fn sample<R>(rng: &mut R, length: usize, amount: usize) -> IndexVec
241where
242 R: Rng + ?Sized,
243{
244 if amount > length {
245 panic!("`amount` of samples must be less than or equal to `length`");
246 }
247 if length > (u32::MAX as usize) {
248 #[cfg(target_pointer_width = "32")]
249 unreachable!();
250
251 #[cfg(target_pointer_width = "64")]
254 return sample_rejection(rng, length as u64, amount as u64);
255 }
256 let amount = amount as u32;
257 let length = length as u32;
258
259 if amount < 163 {
264 const C: [[f32; 2]; 2] = [[1.6, 8.0 / 45.0], [10.0, 70.0 / 9.0]];
265 let j = usize::from(length >= 500_000);
266 let amount_fp = amount as f32;
267 let m4 = C[0][j] * amount_fp;
268 if amount > 11 && (length as f32) < (C[1][j] + m4) * amount_fp {
270 sample_inplace(rng, length, amount)
271 } else {
272 sample_floyd(rng, length, amount)
273 }
274 } else {
275 const C: [f32; 2] = [270.0, 330.0 / 9.0];
276 let j = usize::from(length >= 500_000);
277 if (length as f32) < C[j] * (amount as f32) {
278 sample_inplace(rng, length, amount)
279 } else {
280 sample_rejection(rng, length, amount)
281 }
282 }
283}
284
285#[cfg(feature = "std")]
305pub fn sample_weighted<R, F, X>(
306 rng: &mut R,
307 length: usize,
308 weight: F,
309 amount: usize,
310) -> Result<IndexVec, WeightError>
311where
312 R: Rng + ?Sized,
313 F: Fn(usize) -> X,
314 X: Into<f64>,
315{
316 if length > (u32::MAX as usize) {
317 #[cfg(target_pointer_width = "32")]
318 unreachable!();
319
320 #[cfg(target_pointer_width = "64")]
321 {
322 let amount = amount as u64;
323 let length = length as u64;
324 sample_efraimidis_spirakis(rng, length, weight, amount)
325 }
326 } else {
327 assert!(amount <= u32::MAX as usize);
328 let amount = amount as u32;
329 let length = length as u32;
330 sample_efraimidis_spirakis(rng, length, weight, amount)
331 }
332}
333
334#[cfg(feature = "std")]
352fn sample_efraimidis_spirakis<R, F, X, N>(
353 rng: &mut R,
354 length: N,
355 weight: F,
356 amount: N,
357) -> Result<IndexVec, WeightError>
358where
359 R: Rng + ?Sized,
360 F: Fn(usize) -> X,
361 X: Into<f64>,
362 N: UInt,
363 IndexVec: From<Vec<N>>,
364{
365 use std::{cmp::Ordering, collections::BinaryHeap};
366
367 if amount == N::zero() {
368 return Ok(IndexVec::U32(Vec::new()));
369 }
370
371 struct Element<N> {
372 index: N,
373 key: f64,
374 }
375
376 impl<N> PartialOrd for Element<N> {
377 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
378 Some(self.cmp(other))
379 }
380 }
381
382 impl<N> Ord for Element<N> {
383 fn cmp(&self, other: &Self) -> Ordering {
384 self.key.partial_cmp(&other.key).unwrap().reverse()
387 }
388 }
389
390 impl<N> PartialEq for Element<N> {
391 fn eq(&self, other: &Self) -> bool {
392 self.key == other.key
393 }
394 }
395
396 impl<N> Eq for Element<N> {}
397
398 let mut candidates = BinaryHeap::with_capacity(amount.as_usize());
399 let mut index = N::zero();
400 while index < length && candidates.len() < amount.as_usize() {
401 let weight = weight(index.as_usize()).into();
402 if weight > 0.0 {
403 let key = rng.random::<f64>().ln() / weight;
406 candidates.push(Element { index, key });
407 } else if !(weight >= 0.0) {
408 return Err(WeightError::InvalidWeight);
409 }
410
411 index += N::one();
412 }
413
414 if index < length {
415 let mut x = rng.random::<f64>().ln() / candidates.peek().unwrap().key;
416 while index < length {
417 if !x.is_finite() {
418 return Err(WeightError::InvalidWeight);
419 }
420
421 let weight = weight(index.as_usize()).into();
422 if weight > 0.0 {
423 x -= weight;
424 if x <= 0.0 {
425 let min_candidate = candidates.pop().unwrap();
426 let t = (min_candidate.key * weight).exp();
427 let key = rng.random_range(t..1.0).ln() / weight;
428 candidates.push(Element { index, key });
429
430 x = rng.random::<f64>().ln() / candidates.peek().unwrap().key;
431 }
432 } else if !(weight >= 0.0) {
433 return Err(WeightError::InvalidWeight);
434 }
435
436 index += N::one();
437 }
438 }
439
440 Ok(IndexVec::from(
441 candidates.iter().map(|elt| elt.index).collect(),
442 ))
443}
444
445fn sample_floyd<R>(rng: &mut R, length: u32, amount: u32) -> IndexVec
452where
453 R: Rng + ?Sized,
454{
455 debug_assert!(amount <= length);
459 let mut indices = Vec::with_capacity(amount as usize);
460 for j in length - amount..length {
461 let t = rng.random_range(..=j);
462 if let Some(pos) = indices.iter().position(|&x| x == t) {
463 indices[pos] = j;
464 }
465 indices.push(t);
466 }
467 IndexVec::from(indices)
468}
469
470fn sample_inplace<R>(rng: &mut R, length: u32, amount: u32) -> IndexVec
483where
484 R: Rng + ?Sized,
485{
486 debug_assert!(amount <= length);
487 let mut indices: Vec<u32> = Vec::with_capacity(length as usize);
488 indices.extend(0..length);
489 for i in 0..amount {
490 let j: u32 = rng.random_range(i..length);
491 indices.swap(i as usize, j as usize);
492 }
493 indices.truncate(amount as usize);
494 debug_assert_eq!(indices.len(), amount as usize);
495 IndexVec::from(indices)
496}
497
498trait UInt: Copy + PartialOrd + Ord + PartialEq + Eq + SampleUniform + Hash + AddAssign {
499 fn zero() -> Self;
500 #[cfg_attr(feature = "alloc", allow(dead_code))]
501 fn one() -> Self;
502 fn as_usize(self) -> usize;
503}
504
505impl UInt for u32 {
506 #[inline]
507 fn zero() -> Self {
508 0
509 }
510
511 #[inline]
512 fn one() -> Self {
513 1
514 }
515
516 #[inline]
517 fn as_usize(self) -> usize {
518 self as usize
519 }
520}
521
522#[cfg(target_pointer_width = "64")]
523impl UInt for u64 {
524 #[inline]
525 fn zero() -> Self {
526 0
527 }
528
529 #[inline]
530 fn one() -> Self {
531 1
532 }
533
534 #[inline]
535 fn as_usize(self) -> usize {
536 self as usize
537 }
538}
539
540fn sample_rejection<X: UInt, R>(rng: &mut R, length: X, amount: X) -> IndexVec
550where
551 R: Rng + ?Sized,
552 IndexVec: From<Vec<X>>,
553{
554 debug_assert!(amount < length);
555 #[cfg(feature = "std")]
556 let mut cache = HashSet::with_capacity(amount.as_usize());
557 #[cfg(not(feature = "std"))]
558 let mut cache = BTreeSet::new();
559 let distr = Uniform::new(X::zero(), length).unwrap();
560 let mut indices = Vec::with_capacity(amount.as_usize());
561 for _ in 0..amount.as_usize() {
562 let mut pos = distr.sample(rng);
563 while !cache.insert(pos) {
564 pos = distr.sample(rng);
565 }
566 indices.push(pos);
567 }
568
569 debug_assert_eq!(indices.len(), amount.as_usize());
570 IndexVec::from(indices)
571}
572
573#[cfg(test)]
574mod test {
575 use super::*;
576 use alloc::vec;
577
578 #[test]
579 #[cfg(feature = "serde")]
580 fn test_serialization_index_vec() {
581 let some_index_vec = IndexVec::from(vec![254_u32, 234, 2, 1]);
582 let de_some_index_vec: IndexVec =
583 postcard::from_bytes(&postcard::to_allocvec(&some_index_vec).unwrap()).unwrap();
584 assert_eq!(some_index_vec, de_some_index_vec);
585 }
586
587 #[test]
588 fn test_sample_boundaries() {
589 let mut r = crate::test::rng(404);
590
591 assert_eq!(sample_inplace(&mut r, 0, 0).len(), 0);
592 assert_eq!(sample_inplace(&mut r, 1, 0).len(), 0);
593 assert_eq!(sample_inplace(&mut r, 1, 1).into_vec(), vec![0]);
594
595 assert_eq!(sample_rejection(&mut r, 1u32, 0).len(), 0);
596
597 assert_eq!(sample_floyd(&mut r, 0, 0).len(), 0);
598 assert_eq!(sample_floyd(&mut r, 1, 0).len(), 0);
599 assert_eq!(sample_floyd(&mut r, 1, 1).into_vec(), vec![0]);
600
601 let sum: usize = sample_rejection(&mut r, 1 << 25, 10u32).into_iter().sum();
603 assert!(1 << 25 < sum && sum < (1 << 25) * 25);
604
605 let sum: usize = sample_floyd(&mut r, 1 << 25, 10).into_iter().sum();
606 assert!(1 << 25 < sum && sum < (1 << 25) * 25);
607 }
608
609 #[test]
610 #[cfg_attr(miri, ignore)] fn test_sample_alg() {
612 let seed_rng = crate::test::rng;
613
614 let (length, amount): (usize, usize) = (100, 50);
620 let v1 = sample(&mut seed_rng(420), length, amount);
621 let v2 = sample_inplace(&mut seed_rng(420), length as u32, amount as u32);
622 assert!(v1.iter().all(|e| e < length));
623 assert_eq!(v1, v2);
624
625 let v3 = sample_floyd(&mut seed_rng(420), length as u32, amount as u32);
627 assert!(v1 != v3);
628
629 let (length, amount): (usize, usize) = (1 << 20, 50);
631 let v1 = sample(&mut seed_rng(421), length, amount);
632 let v2 = sample_floyd(&mut seed_rng(421), length as u32, amount as u32);
633 assert!(v1.iter().all(|e| e < length));
634 assert_eq!(v1, v2);
635
636 let (length, amount): (usize, usize) = (1 << 20, 600);
638 let v1 = sample(&mut seed_rng(422), length, amount);
639 let v2 = sample_rejection(&mut seed_rng(422), length as u32, amount as u32);
640 assert!(v1.iter().all(|e| e < length));
641 assert_eq!(v1, v2);
642 }
643
644 #[cfg(feature = "std")]
645 #[test]
646 fn test_sample_weighted() {
647 let seed_rng = crate::test::rng;
648 for &(amount, len) in &[(0, 10), (5, 10), (9, 10)] {
649 let v = sample_weighted(&mut seed_rng(423), len, |i| i as f64, amount).unwrap();
650 match v {
651 IndexVec::U32(mut indices) => {
652 assert_eq!(indices.len(), amount);
653 indices.sort_unstable();
654 indices.dedup();
655 assert_eq!(indices.len(), amount);
656 for &i in &indices {
657 assert!((i as usize) < len);
658 }
659 }
660 #[cfg(target_pointer_width = "64")]
661 _ => panic!("expected `IndexVec::U32`"),
662 }
663 }
664
665 let r = sample_weighted(&mut seed_rng(423), 10, |i| i as f64, 10);
666 assert_eq!(r.unwrap().len(), 9);
667 }
668
669 #[cfg(feature = "std")]
670 #[test]
671 fn test_sample_weighted_infinities() {
672 let mut rng = crate::test::rng(1351);
673
674 for _ in 0..10 {
675 let result = sample_weighted(&mut rng, 5, |i| 2.0 / (2.0 - (i as f64)).abs(), 2);
676 assert!(result.is_ok());
677 assert!(result.unwrap().iter().any(|i| i == 2));
679 }
680
681 let weights = [1.0, 0.0, f32::INFINITY, 2.5, f32::INFINITY];
683 let result = sample_weighted(&mut rng, weights.len(), |i| weights[i], 2);
684 assert!(result.is_ok());
685 let mut results = result.unwrap().into_vec();
686 results.sort();
687 assert_eq!(results, [2, 4]);
688
689 assert!(sample_weighted(&mut rng, weights.len(), |i| weights[i], 1).is_err());
691 }
692
693 #[test]
694 fn value_stability_sample() {
695 let do_test = |length, amount, values: &[u32]| {
696 let mut buf = [0u32; 8];
697 let mut rng = crate::test::rng(410);
698
699 let res = sample(&mut rng, length, amount);
700 let len = res.len().min(buf.len());
701 for (x, y) in res.into_iter().zip(buf.iter_mut()) {
702 *y = x as u32;
703 }
704 assert_eq!(
705 &buf[0..len],
706 values,
707 "failed sampling {}, {}",
708 length,
709 amount
710 );
711 };
712
713 do_test(10, 6, &[0, 9, 5, 4, 6, 8]); do_test(25, 10, &[24, 20, 19, 9, 22, 16, 0, 14]); do_test(300, 8, &[30, 283, 243, 150, 218, 240, 1, 189]); do_test(300, 80, &[31, 289, 248, 154, 221, 243, 7, 192]); do_test(300, 180, &[31, 289, 248, 154, 221, 243, 7, 192]); do_test(
720 1_000_000,
721 8,
722 &[103717, 963485, 826422, 509101, 736394, 807035, 5327, 632573],
723 ); do_test(
725 1_000_000,
726 180,
727 &[103718, 963490, 826426, 509103, 736396, 807036, 5327, 632573],
728 ); }
730}