Skip to main content

del_msh_cpu/
alias_table.rs

1use rand::{Rng, RngExt};
2use std::collections::VecDeque;
3
4/// 離散分布を O(1) でサンプリングする Alias Table
5pub struct AliasTable {
6    pub prob: Vec<f32>,  // 各スロットのメイン確率 (0..=1)
7    pub alias: Vec<u32>, // メインが外れたときの代替インデックス
8    pub sum_w: f64,      // 元の重みの総和(pdf計算に使う)
9}
10
11impl AliasTable {
12    /// 重み配列から AliasTable を構築する
13    /// weights[i] >= 0 を想定
14    pub fn new(weights: &[f32]) -> Self {
15        let n = weights.len();
16        assert!(n > 0);
17
18        // 総和
19        let sum_w: f64 = weights.iter().map(|&w| w.max(0.0) as f64).sum();
20
21        // 全てゼロの場合は一様分布にする
22        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        // 正規化された確率を N 倍したもの
33        // p[i] = weights[i] / sum_w * N
34        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        // Vose のアルゴリズム
55        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; // p[l] -= (1 - p[s])
60
61            if p[l] < 1.0 {
62                small.push_back(l);
63            } else {
64                large.push_back(l);
65            }
66        }
67
68        // 残りは全て 1 に丸める
69        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    /// 1 サンプル取得(添字を返す)
78    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(); // [0,1)
84        if r < self.prob[i] {
85            i
86        } else {
87            self.alias[i] as usize
88        }
89    }
90}