Skip to main content

zsfm_tabfm/ensemble/
pyrandom.rs

1//! A bit-compatible port of CPython's `random.Random` (Mersenne Twister, MT19937), so ensemble
2//! member generation can be validated numerically against the real `TabFMClassifier`/
3//! `TabFMRegressor` sklearn wrapper (which seeds `random.Random(random_state)` and calls
4//! `.sample()`/`.shuffle()` in a specific order — see `config_gen.rs`).
5//!
6//! Ported from CPython's `Modules/_randommodule.c` (MT19937 core, `init_by_array` seeding) and
7//! `Lib/random.py` (`_randbelow`, `sample`, `shuffle`). Reference sequences used in the unit
8//! tests below were generated with the project's own `.venv/bin/python3`.
9
10const 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    /// CPython seeds every integer via `init_by_array`, never bare `init_genrand` — the seed
33    /// integer is first split into little-endian 32-bit words (`key`).
34    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    /// Seed from a small non-negative integer (as CPython does for `random.Random(42)` etc).
70    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    /// CPython's `getrandbits(k)`: words filled least-significant-first, 32 bits at a time.
112    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
129/// A CPython-compatible `random.Random` instance: MT19937 state plus the pure-Python wrapper
130/// logic (`_randbelow`, `sample`, `shuffle`) from `Lib/random.py`.
131pub 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    /// `Random._randbelow_with_getrandbits`: rejection-sampled `getrandbits`.
145    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    /// `Random.sample(range(n), k)` restricted to integer populations `0..n` (the only case
159    /// needed here — feature/class/row indices), returning the sampled indices in draw order.
160    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    /// `Random.sample(population, k)` for an arbitrary slice, via `sample_indices`.
189    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    /// `Random.shuffle(x)`: in-place Fisher-Yates via `_randbelow`.
197    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        // python3 -c "import random; print(random.Random(42).sample(range(10), 10))"
216        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        // python3 -c "import random; print(random.Random(42).sample(range(5), 3))"
223        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        // python3 -c "import random; l=list(range(8)); random.Random(42).shuffle(l); print(l)"
230        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        // python3 -c "import random; print(random.Random(7).sample(range(100), 5))"
239        // k=5 (not >5) so setsize=21; n=100 > 21, exercises the rejection-set branch.
240        let mut r = PyRandom::new(7);
241        assert_eq!(r.sample_indices(100, 5), vec![41, 19, 50, 83, 6]);
242    }
243}