const MAXBIT: u32 = 30;
pub const MAX_SOBOL_DIM: usize = 7;
const SOBOL_TABLE: [(u32, u32, [u32; 4]); MAX_SOBOL_DIM - 1] = [
(1, 0, [1, 0, 0, 0]),
(2, 1, [1, 3, 0, 0]),
(3, 1, [1, 3, 1, 0]),
(3, 2, [1, 1, 1, 0]),
(4, 1, [1, 1, 3, 3]),
(4, 4, [1, 3, 5, 13]),
];
const PRIMES: [u64; 32] = [
2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41, 43, 47, 53, 59, 61, 67, 71, 73, 79, 83, 89, 97,
101, 103, 107, 109, 113, 127, 131,
];
pub fn halton_radical_inverse(index: usize, base: u64) -> f64 {
let mut result = 0.0_f64;
let mut f = 1.0 / base as f64;
let mut i = index as u64;
while i > 0 {
result += f * (i % base) as f64;
i /= base;
f /= base as f64;
}
result
}
fn build_direction_numbers(mdeg: u32, ip: u32, iv: [u32; 4]) -> Vec<u32> {
let maxbit = MAXBIT as usize;
let deg = mdeg as usize;
let mut v = vec![0u32; maxbit + 1];
for i in 1..=deg {
v[i] = iv[i - 1] << (MAXBIT - i as u32);
}
let nbits = deg.saturating_sub(1);
let mut bits = vec![0u32; nbits];
for (k, bit) in bits.iter_mut().enumerate() {
*bit = (ip >> (nbits - 1 - k)) & 1;
}
for i in (deg + 1)..=maxbit {
let mut vi = v[i - deg];
vi ^= v[i - deg] >> mdeg;
for (k, &bit) in bits.iter().enumerate() {
if bit != 0 {
vi ^= v[i - (k + 1)];
}
}
v[i] = vi;
}
v
}
pub struct SobolGenerator {
dim: usize,
direction_numbers: Vec<Vec<u32>>,
x: Vec<u32>,
count: usize,
}
impl SobolGenerator {
pub fn new(dim: usize) -> Self {
let direction_numbers = (0..dim)
.map(|d| {
if d == 0 {
let mut v = vec![0u32; MAXBIT as usize + 1];
for (i, slot) in v.iter_mut().enumerate().skip(1) {
*slot = 1u32 << (MAXBIT - i as u32);
}
v
} else if d < MAX_SOBOL_DIM {
let (mdeg, ip, iv) = SOBOL_TABLE[d - 1];
build_direction_numbers(mdeg, ip, iv)
} else {
Vec::new() }
})
.collect();
Self {
dim,
direction_numbers,
x: vec![0u32; dim],
count: 0,
}
}
pub fn next_point(&mut self) -> Vec<f64> {
let n = self.count;
self.count += 1;
let point: Vec<f64> = (0..self.dim)
.map(|d| {
if d < MAX_SOBOL_DIM {
self.x[d] as f64 / (1u64 << MAXBIT) as f64
} else {
let base = PRIMES[d % PRIMES.len()];
halton_radical_inverse(n + 1, base)
}
})
.collect();
let c = (n as u32).trailing_ones() as usize + 1;
for d in 0..self.dim.min(MAX_SOBOL_DIM) {
if let Some(&vc) = self.direction_numbers[d].get(c) {
self.x[d] ^= vc;
}
}
point
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_halton_radical_inverse_base2_matches_van_der_corput() {
assert!((halton_radical_inverse(1, 2) - 0.5).abs() < 1e-12);
assert!((halton_radical_inverse(2, 2) - 0.25).abs() < 1e-12);
assert!((halton_radical_inverse(3, 2) - 0.75).abs() < 1e-12);
assert!((halton_radical_inverse(4, 2) - 0.125).abs() < 1e-12);
}
#[test]
fn test_halton_radical_inverse_stays_in_unit_interval() {
for base in [2, 3, 5, 7, 11] {
for index in 1..100 {
let v = halton_radical_inverse(index, base);
assert!((0.0..1.0).contains(&v), "base={base} index={index} v={v}");
}
}
}
#[test]
fn test_sobol_matches_scipy_reference_7d() {
let expected: [[f64; 7]; 8] = [
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
[0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5],
[0.75, 0.25, 0.25, 0.25, 0.75, 0.75, 0.25],
[0.25, 0.75, 0.75, 0.75, 0.25, 0.25, 0.75],
[0.375, 0.375, 0.625, 0.875, 0.375, 0.125, 0.375],
[0.875, 0.875, 0.125, 0.375, 0.875, 0.625, 0.875],
[0.625, 0.125, 0.875, 0.625, 0.625, 0.875, 0.125],
[0.125, 0.625, 0.375, 0.125, 0.125, 0.375, 0.625],
];
let mut gen = SobolGenerator::new(7);
for row in expected.iter() {
let p = gen.next_point();
for (got, &want) in p.iter().zip(row.iter()) {
assert!((got - want).abs() < 1e-12, "got {got}, want {want}");
}
}
}
#[test]
fn test_sobol_points_stay_in_unit_interval_and_are_non_constant() {
let mut gen = SobolGenerator::new(5);
let mut first_coords = Vec::new();
for _ in 0..50 {
let p = gen.next_point();
assert_eq!(p.len(), 5);
for &v in &p {
assert!((0.0..1.0).contains(&v));
}
first_coords.push(p[0]);
}
assert!(first_coords.iter().any(|&v| v != first_coords[0]));
}
#[test]
fn test_sobol_fallback_dimension_beyond_table_is_still_low_discrepancy() {
let dim = MAX_SOBOL_DIM + 2;
let mut gen = SobolGenerator::new(dim);
let mut last_col = Vec::new();
for _ in 0..20 {
let p = gen.next_point();
for &v in &p {
assert!((0.0..1.0).contains(&v));
}
last_col.push(p[dim - 1]);
}
assert!(last_col.iter().any(|&v| v != last_col[0]));
}
}