zsfm_tabfm/ensemble/
pyrandom.rs1const N: usize = 624;
11const M: usize = 397;
12const MATRIX_A: u32 = 0x9908_b0df;
13const UPPER_MASK: u32 = 0x8000_0000;
14const LOWER_MASK: u32 = 0x7fff_ffff;
15
16pub struct Mt19937 {
17 mt: [u32; N],
18 index: usize,
19}
20
21impl Mt19937 {
22 fn init_genrand(seed: u32) -> Self {
23 let mut mt = [0u32; N];
24 mt[0] = seed;
25 for i in 1..N {
26 mt[i] = (1_812_433_253u32.wrapping_mul(mt[i - 1] ^ (mt[i - 1] >> 30)))
27 .wrapping_add(i as u32);
28 }
29 Mt19937 { mt, index: N }
30 }
31
32 pub fn from_seed_key(key: &[u32]) -> Self {
35 let mut rng = Self::init_genrand(19_650_218);
36 let key = if key.is_empty() { vec![0u32] } else { key.to_vec() };
37 let key_length = key.len();
38 let mut i = 1usize;
39 let mut j = 0usize;
40 for _ in 0..N.max(key_length) {
41 let prev = rng.mt[i - 1];
42 rng.mt[i] = (rng.mt[i] ^ ((prev ^ (prev >> 30)).wrapping_mul(1_664_525)))
43 .wrapping_add(key[j])
44 .wrapping_add(j as u32);
45 i += 1;
46 j += 1;
47 if i >= N {
48 rng.mt[0] = rng.mt[N - 1];
49 i = 1;
50 }
51 if j >= key_length {
52 j = 0;
53 }
54 }
55 for _ in 0..N - 1 {
56 let prev = rng.mt[i - 1];
57 rng.mt[i] = (rng.mt[i] ^ ((prev ^ (prev >> 30)).wrapping_mul(1_566_083_941)))
58 .wrapping_sub(i as u32);
59 i += 1;
60 if i >= N {
61 rng.mt[0] = rng.mt[N - 1];
62 i = 1;
63 }
64 }
65 rng.mt[0] = 0x8000_0000;
66 rng
67 }
68
69 pub fn from_u64_seed(seed: u64) -> Self {
71 if seed == 0 {
72 return Self::from_seed_key(&[]);
73 }
74 let mut key = Vec::new();
75 let mut s = seed;
76 while s > 0 {
77 key.push((s & 0xffff_ffff) as u32);
78 s >>= 32;
79 }
80 Self::from_seed_key(&key)
81 }
82
83 fn regenerate(&mut self) {
84 let mag01 = [0u32, MATRIX_A];
85 for kk in 0..N - M {
86 let y = (self.mt[kk] & UPPER_MASK) | (self.mt[kk + 1] & LOWER_MASK);
87 self.mt[kk] = self.mt[kk + M] ^ (y >> 1) ^ mag01[(y & 1) as usize];
88 }
89 for kk in N - M..N - 1 {
90 let y = (self.mt[kk] & UPPER_MASK) | (self.mt[kk + 1] & LOWER_MASK);
91 self.mt[kk] = self.mt[kk + M - N] ^ (y >> 1) ^ mag01[(y & 1) as usize];
92 }
93 let y = (self.mt[N - 1] & UPPER_MASK) | (self.mt[0] & LOWER_MASK);
94 self.mt[N - 1] = self.mt[M - 1] ^ (y >> 1) ^ mag01[(y & 1) as usize];
95 self.index = 0;
96 }
97
98 pub fn next_u32(&mut self) -> u32 {
99 if self.index >= N {
100 self.regenerate();
101 }
102 let mut y = self.mt[self.index];
103 self.index += 1;
104 y ^= y >> 11;
105 y ^= (y << 7) & 0x9d2c_5680;
106 y ^= (y << 15) & 0xefc6_0000;
107 y ^= y >> 18;
108 y
109 }
110
111 pub fn getrandbits(&mut self, k: u32) -> u64 {
113 if k <= 32 {
114 return (self.next_u32() >> (32 - k)) as u64;
115 }
116 let words = (k - 1) / 32 + 1;
117 let mut result: u64 = 0;
118 for w in 0..words {
119 let mut r = self.next_u32();
120 if w == words - 1 {
121 r >>= 32 * words - k;
122 }
123 result |= (r as u64) << (32 * w);
124 }
125 result
126 }
127}
128
129pub struct PyRandom {
132 mt: Mt19937,
133}
134
135impl PyRandom {
136 pub fn new(seed: u64) -> Self {
137 PyRandom { mt: Mt19937::from_u64_seed(seed) }
138 }
139
140 fn bit_length(n: u64) -> u32 {
141 if n == 0 { 0 } else { 64 - n.leading_zeros() }
142 }
143
144 pub fn randbelow(&mut self, n: u64) -> u64 {
146 if n == 0 {
147 return 0;
148 }
149 let k = Self::bit_length(n);
150 loop {
151 let r = self.mt.getrandbits(k);
152 if r < n {
153 return r;
154 }
155 }
156 }
157
158 pub fn sample_indices(&mut self, n: usize, k: usize) -> Vec<usize> {
161 assert!(k <= n, "sample larger than population");
162 let mut setsize: f64 = 21.0;
163 if k > 5 {
164 setsize += 4f64.powf((3.0 * k as f64).log(4.0).ceil());
165 }
166 let mut result = vec![0usize; k];
167 if (n as f64) <= setsize {
168 let mut pool: Vec<usize> = (0..n).collect();
169 for i in 0..k {
170 let j = self.randbelow((n - i) as u64) as usize;
171 result[i] = pool[j];
172 pool[j] = pool[n - i - 1];
173 }
174 } else {
175 let mut selected = std::collections::HashSet::new();
176 for i in 0..k {
177 let mut j = self.randbelow(n as u64) as usize;
178 while selected.contains(&j) {
179 j = self.randbelow(n as u64) as usize;
180 }
181 selected.insert(j);
182 result[i] = j;
183 }
184 }
185 result
186 }
187
188 pub fn sample<T: Clone>(&mut self, population: &[T], k: usize) -> Vec<T> {
190 self.sample_indices(population.len(), k)
191 .into_iter()
192 .map(|i| population[i].clone())
193 .collect()
194 }
195
196 pub fn shuffle<T>(&mut self, x: &mut [T]) {
198 let len = x.len();
199 if len < 2 {
200 return;
201 }
202 for i in (1..len).rev() {
203 let j = self.randbelow((i + 1) as u64) as usize;
204 x.swap(i, j);
205 }
206 }
207}
208
209#[cfg(test)]
210mod tests {
211 use super::*;
212
213 #[test]
214 fn test_sample_10_of_10() {
215 let mut r = PyRandom::new(42);
217 assert_eq!(r.sample_indices(10, 10), vec![1, 0, 4, 9, 6, 5, 8, 2, 3, 7]);
218 }
219
220 #[test]
221 fn test_sample_3_of_5() {
222 let mut r = PyRandom::new(42);
224 assert_eq!(r.sample_indices(5, 3), vec![0, 4, 2]);
225 }
226
227 #[test]
228 fn test_shuffle_8() {
229 let mut r = PyRandom::new(42);
231 let mut v: Vec<usize> = (0..8).collect();
232 r.shuffle(&mut v);
233 assert_eq!(v, vec![3, 4, 6, 7, 2, 5, 0, 1]);
234 }
235
236 #[test]
237 fn test_sample_5_of_100_rejection_branch() {
238 let mut r = PyRandom::new(7);
241 assert_eq!(r.sample_indices(100, 5), vec![41, 19, 50, 83, 6]);
242 }
243}