del_msh_cpu/
alias_table.rs1use rand::{Rng, RngExt};
2use std::collections::VecDeque;
3
4pub struct AliasTable {
6 pub prob: Vec<f32>, pub alias: Vec<u32>, pub sum_w: f64, }
10
11impl AliasTable {
12 pub fn new(weights: &[f32]) -> Self {
15 let n = weights.len();
16 assert!(n > 0);
17
18 let sum_w: f64 = weights.iter().map(|&w| w.max(0.0) as f64).sum();
20
21 if sum_w == 0.0 {
23 let prob = vec![1.0f32; n];
24 let alias = (0..n as u32).collect();
25 return AliasTable {
26 prob,
27 alias,
28 sum_w: n as f64,
29 };
30 }
31
32 let n_f = n as f64;
35 let mut p: Vec<f64> = weights
36 .iter()
37 .map(|&w| (w.max(0.0) as f64) * n_f / sum_w)
38 .collect();
39
40 let mut small = VecDeque::new();
41 let mut large = VecDeque::new();
42
43 for (i, &pi) in p.iter().enumerate() {
44 if pi < 1.0 {
45 small.push_back(i);
46 } else {
47 large.push_back(i);
48 }
49 }
50
51 let mut prob = vec![0.0f32; n];
52 let mut alias = vec![0u32; n];
53
54 while let (Some(s), Some(l)) = (small.pop_front(), large.pop_front()) {
56 prob[s] = p[s] as f32;
57 alias[s] = l as u32;
58
59 p[l] = (p[l] + p[s]) - 1.0; if p[l] < 1.0 {
62 small.push_back(l);
63 } else {
64 large.push_back(l);
65 }
66 }
67
68 for i in large.into_iter().chain(small) {
70 prob[i] = 1.0;
71 alias[i] = i as u32;
72 }
73
74 AliasTable { prob, alias, sum_w }
75 }
76
77 pub fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> usize {
79 let n = self.prob.len();
80 debug_assert!(n > 0);
81
82 let i = rng.random_range(0..n);
83 let r: f32 = rng.random(); if r < self.prob[i] {
85 i
86 } else {
87 self.alias[i] as usize
88 }
89 }
90}