use crate::error::{StatsError, StatsResult};
use scirs2_core::ndarray::{Array1, Array2};
use scirs2_core::numeric::{Float, FromPrimitive, One, Zero};
use scirs2_core::random::{rngs::StdRng, Rng, RngExt, SeedableRng};
use scirs2_core::{parallel_ops::*, simd_ops::SimdUnifiedOps, validation::*};
use std::marker::PhantomData;
const FAURE_DIGITS: usize = 32;
const QUALITY_SAMPLE_CAP: usize = 200;
const WRAPAROUND_DIM_CAP: usize = 20;
const DIAPHONY_DIM_CAP: usize = 16;
const DIAPHONY_FREQ_MAX: i64 = 2;
pub struct EnhancedQMCGenerator<F> {
pub sequence_type: EnhancedSequenceType,
pub dimension: usize,
pub config: EnhancedQMCConfig,
pub state: QMCGeneratorState,
sobol_direction_numbers: Vec<Vec<u64>>,
niederreiter_generating_matrices: Vec<Array2<u32>>,
faure_base: u64,
faure_pascal_matrix: Array2<u64>,
_phantom: PhantomData<F>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum EnhancedSequenceType {
SobolAdvanced {
owen_scrambling: bool,
digital_shift: bool,
nested_scrambling: bool,
},
Niederreiter {
base_strategy: BaseSelectionStrategy,
matrix_optimization: bool,
},
FaureImproved {
permutation_optimization: bool,
radical_inverse_improvements: bool,
},
DigitalNet {
net_params: DigitalNetParams,
construction_method: NetConstructionMethod,
},
Hybrid {
primary: Box<EnhancedSequenceType>,
secondary: Box<EnhancedSequenceType>,
combination: HybridCombinationStrategy,
},
}
#[derive(Debug, Clone, PartialEq)]
pub enum BaseSelectionStrategy {
FirstPrimes,
OptimizedPrimes,
PrimePowers,
Automatic,
}
#[derive(Debug, Clone, PartialEq)]
pub struct DigitalNetParams {
pub t: usize,
pub m: usize,
pub s: usize,
pub base: usize,
}
#[derive(Debug, Clone, PartialEq)]
pub enum NetConstructionMethod {
Sobol,
NiederreiterXing,
PolynomialLattice,
FiniteField,
}
#[derive(Debug, Clone, PartialEq)]
pub enum HybridCombinationStrategy {
Interleave,
Weighted(f64),
DimensionAlternation,
Adaptive,
}
#[derive(Debug, Clone)]
pub struct EnhancedQMCConfig {
pub parallel: bool,
pub chunksize: usize,
pub seed: Option<u64>,
pub use_simd: bool,
pub quality_threshold: f64,
pub max_assessment_length: usize,
pub adaptive_refinement: bool,
}
impl Default for EnhancedQMCConfig {
fn default() -> Self {
Self {
parallel: true,
chunksize: 1000,
seed: None,
use_simd: true,
quality_threshold: 1e-3,
max_assessment_length: 10000,
adaptive_refinement: false,
}
}
}
#[derive(Debug, Clone)]
pub struct QMCGeneratorState {
pub current_index: usize,
pub scrambling_matrices: Option<Vec<Array2<u32>>>,
pub digital_shifts: Option<Vec<Array1<u32>>>,
pub quality_metrics: QualityMetrics,
}
#[derive(Debug, Clone, Default)]
pub struct QualityMetrics {
pub star_discrepancy: f64,
pub wraparound_discrepancy: f64,
pub diaphony: f64,
pub figure_of_merit: f64,
}
impl<F> EnhancedQMCGenerator<F>
where
F: Float + Zero + One + Copy + Send + Sync + SimdUnifiedOps + FromPrimitive + std::fmt::Display,
{
pub fn new(
sequence_type: EnhancedSequenceType,
dimension: usize,
config: EnhancedQMCConfig,
) -> StatsResult<Self> {
check_positive(dimension, "dimension")?;
if dimension > 1000 {
return Err(StatsError::InvalidArgument(
"Dimension cannot exceed 1000 for enhanced QMC sequences".to_string(),
));
}
let state = QMCGeneratorState {
current_index: 0,
scrambling_matrices: None,
digital_shifts: None,
quality_metrics: QualityMetrics::default(),
};
let sobol_direction_numbers =
crate::qmc::advanced::AdvancedQMCGenerator::load_joe_kuo_direction_numbers(dimension)
.map_err(|e| {
StatsError::ComputationError(format!(
"Failed to initialize Sobol direction numbers: {e}"
))
})?;
let niederreiter_matrix_optimization = match &sequence_type {
EnhancedSequenceType::Niederreiter {
matrix_optimization,
..
} => *matrix_optimization,
_ => true,
};
let niederreiter_generating_matrices =
crate::qmc::advanced::AdvancedQMCGenerator::generate_niederreiter_matrices(
dimension,
niederreiter_matrix_optimization,
)
.map_err(|e| {
StatsError::ComputationError(format!(
"Failed to initialize Niederreiter generating matrices: {e}"
))
})?;
let faure_base = u64::from(Self::smallest_prime_geq(dimension as u32));
let faure_pascal_matrix = Self::build_faure_pascal_matrix(faure_base);
let mut generator = Self {
sequence_type,
dimension,
config,
state,
sobol_direction_numbers,
niederreiter_generating_matrices,
faure_base,
faure_pascal_matrix,
_phantom: PhantomData,
};
generator.initialize_randomization()?;
Ok(generator)
}
pub fn generate(&mut self, n: usize) -> StatsResult<Array2<F>> {
check_positive(n, "n")?;
if self.config.parallel && n >= self.config.chunksize {
self.generate_parallel(n)
} else {
self.generate_sequential(n)
}
}
fn generate_parallel(&mut self, n: usize) -> StatsResult<Array2<F>> {
let chunksize = self.config.chunksize;
let num_chunks = n.div_ceil(chunksize);
let chunks = parallel_map_result(
(0..num_chunks).collect::<Vec<_>>().as_slice(),
|&chunk_idx| {
let start = chunk_idx * chunksize;
let end = (start + chunksize).min(n);
let chunksize = end - start;
self.generate_chunk(start, chunksize)
},
)?;
let mut result = Array2::zeros((n, self.dimension));
let mut row_idx = 0;
for chunk in chunks {
let chunk = chunk;
let chunk_rows = chunk.nrows();
result
.slice_mut(scirs2_core::ndarray::s![row_idx..row_idx + chunk_rows, ..])
.assign(&chunk);
row_idx += chunk_rows;
}
if n <= self.config.max_assessment_length {
self.assess_quality(&result)?;
}
Ok(result)
}
fn generate_sequential(&mut self, n: usize) -> StatsResult<Array2<F>> {
let mut result = Array2::zeros((n, self.dimension));
for i in 0..n {
let point = self.next_point()?;
result.row_mut(i).assign(&point);
}
if n <= self.config.max_assessment_length {
self.assess_quality(&result)?;
}
Ok(result)
}
fn generate_chunk(&self, start_index: usize, chunksize: usize) -> StatsResult<Array2<F>> {
let mut chunk = Array2::zeros((chunksize, self.dimension));
for i in 0..chunksize {
let _index = start_index + i;
let point = self.compute_point_at_index(_index)?;
chunk.row_mut(i).assign(&point);
}
Ok(chunk)
}
fn next_point(&mut self) -> StatsResult<Array1<F>> {
let point = self.compute_point_at_index(self.state.current_index)?;
self.state.current_index += 1;
Ok(point)
}
fn compute_point_at_index(&self, index: usize) -> StatsResult<Array1<F>> {
self.compute_point_for_type(index, &self.sequence_type)
}
fn compute_point_for_type(
&self,
index: usize,
seq_type: &EnhancedSequenceType,
) -> StatsResult<Array1<F>> {
match seq_type {
EnhancedSequenceType::SobolAdvanced {
owen_scrambling,
digital_shift,
nested_scrambling,
} => self.compute_sobol_advanced(
index,
*owen_scrambling,
*digital_shift,
*nested_scrambling,
),
EnhancedSequenceType::Niederreiter {
base_strategy,
matrix_optimization,
} => self.compute_niederreiter_enhanced(index, base_strategy, *matrix_optimization),
EnhancedSequenceType::FaureImproved {
permutation_optimization,
radical_inverse_improvements,
} => self.compute_faure_improved(
index,
*permutation_optimization,
*radical_inverse_improvements,
),
EnhancedSequenceType::DigitalNet {
net_params,
construction_method,
} => self.compute_digital_net(index, net_params, construction_method),
EnhancedSequenceType::Hybrid {
primary,
secondary,
combination,
} => self.compute_hybrid_sequence(index, primary, secondary, combination),
}
}
fn compute_sobol_advanced(
&self,
index: usize,
owen_scrambling: bool,
digital_shift: bool,
_nested_scrambling: bool,
) -> StatsResult<Array1<F>> {
let mut point = Array1::zeros(self.dimension);
let gray_code = index ^ (index >> 1);
for dim in 0..self.dimension {
let dir_nums = &self.sobol_direction_numbers[dim];
let mut result_64 = 0u64;
for (bit, &dn) in dir_nums.iter().enumerate().take(32) {
if (gray_code >> bit) & 1 == 1 {
result_64 ^= dn;
}
}
let mut result = (result_64 >> 32) as u32;
if owen_scrambling {
if let Some(ref matrices) = self.state.scrambling_matrices {
if dim < matrices.len() {
result = self.apply_owen_scrambling(result, &matrices[dim]);
}
}
}
if digital_shift {
if let Some(ref shifts) = self.state.digital_shifts {
if dim < shifts.len() {
result ^= shifts[dim][0]; }
}
}
point[dim] = F::from(result as f64 / (1u64 << 32) as f64).expect("Operation failed");
}
Ok(point)
}
fn compute_niederreiter_enhanced(
&self,
index: usize,
_base_strategy: &BaseSelectionStrategy,
_matrix_optimization: bool,
) -> StatsResult<Array1<F>> {
let raw = crate::qmc::advanced::AdvancedQMCGenerator::niederreiter_point_from_matrices(
self.dimension,
index,
&self.niederreiter_generating_matrices,
);
let mut point = Array1::zeros(self.dimension);
for dim in 0..self.dimension {
point[dim] = F::from(raw[dim]).expect("Operation failed");
}
Ok(point)
}
fn compute_faure_improved(
&self,
index: usize,
_permutation_optimization: bool,
_radical_inverse_improvements: bool,
) -> StatsResult<Array1<F>> {
let base = self.faure_base;
let mut digits = [0u64; FAURE_DIGITS];
let mut rem = index as u64;
for slot in digits.iter_mut() {
*slot = rem % base;
rem /= base;
}
let mut point = Array1::zeros(self.dimension);
let mut c = digits;
for dim in 0..self.dimension {
if dim > 0 {
let mut next = [0u64; FAURE_DIGITS];
for (r, slot) in next.iter_mut().enumerate() {
let mut acc = 0u64;
for l in r..FAURE_DIGITS {
acc += self.faure_pascal_matrix[[r, l]] * c[l];
}
*slot = acc % base;
}
c = next;
}
let mut value = 0.0f64;
let mut fraction = 1.0 / base as f64;
for &digit in c.iter() {
value += digit as f64 * fraction;
fraction /= base as f64;
}
point[dim] = F::from(value).expect("Operation failed");
}
Ok(point)
}
fn compute_digital_net(
&self,
index: usize,
net_params: &DigitalNetParams,
construction_method: &NetConstructionMethod,
) -> StatsResult<Array1<F>> {
if net_params.base != 2 {
return Err(StatsError::NotImplementedError(format!(
"DigitalNet with base {} is not implemented: this crate's digital-net \
constructions (Sobol, NiederreiterXing) are base-2 only",
net_params.base
)));
}
match construction_method {
NetConstructionMethod::Sobol => self.compute_sobol_advanced(index, false, false, false),
NetConstructionMethod::NiederreiterXing => {
self.compute_niederreiter_enhanced(index, &BaseSelectionStrategy::Automatic, true)
}
NetConstructionMethod::PolynomialLattice | NetConstructionMethod::FiniteField => {
Err(StatsError::NotImplementedError(format!(
"DigitalNet construction method {construction_method:?} is not implemented: \
it requires polynomial-ring/finite-field arithmetic that this crate's QMC \
module does not yet provide; use NetConstructionMethod::Sobol or \
NetConstructionMethod::NiederreiterXing instead"
)))
}
}
}
fn compute_hybrid_sequence(
&self,
index: usize,
primary: &EnhancedSequenceType,
secondary: &EnhancedSequenceType,
combination: &HybridCombinationStrategy,
) -> StatsResult<Array1<F>> {
let p = self.compute_point_for_type(index, primary)?;
let s = self.compute_point_for_type(index, secondary)?;
let mut point = Array1::zeros(self.dimension);
match combination {
HybridCombinationStrategy::Interleave => {
let use_primary = index.is_multiple_of(2);
for dim in 0..self.dimension {
point[dim] = if use_primary { p[dim] } else { s[dim] };
}
}
HybridCombinationStrategy::Weighted(w) => {
let w = F::from(w.clamp(0.0, 1.0)).expect("Operation failed");
let one_minus_w = F::one() - w;
for dim in 0..self.dimension {
point[dim] = w * p[dim] + one_minus_w * s[dim];
}
}
HybridCombinationStrategy::DimensionAlternation => {
for dim in 0..self.dimension {
point[dim] = if dim.is_multiple_of(2) {
p[dim]
} else {
s[dim]
};
}
}
HybridCombinationStrategy::Adaptive => {
let half = F::from(0.5).expect("Operation failed");
for dim in 0..self.dimension {
point[dim] = half * p[dim] + half * s[dim];
}
}
}
Ok(point)
}
fn initialize_randomization(&mut self) -> StatsResult<()> {
let mut rng = match self.config.seed {
Some(seed) => StdRng::seed_from_u64(seed),
None => StdRng::from_rng(&mut scirs2_core::random::thread_rng()),
};
if self.needs_scrambling() {
let mut matrices = Vec::with_capacity(self.dimension);
for _ in 0..self.dimension {
matrices.push(self.generate_scrambling_matrix(&mut rng)?);
}
self.state.scrambling_matrices = Some(matrices);
}
if self.needs_digital_shift() {
let mut shifts = Vec::with_capacity(self.dimension);
for _ in 0..self.dimension {
let shift = Array1::from_shape_fn(32, |_| rng.random::<u32>());
shifts.push(shift);
}
self.state.digital_shifts = Some(shifts);
}
Ok(())
}
fn needs_scrambling(&self) -> bool {
match &self.sequence_type {
EnhancedSequenceType::SobolAdvanced {
owen_scrambling, ..
} => *owen_scrambling,
_ => false,
}
}
fn needs_digital_shift(&self) -> bool {
match &self.sequence_type {
EnhancedSequenceType::SobolAdvanced { digital_shift, .. } => *digital_shift,
_ => false,
}
}
fn generate_scrambling_matrix<R: Rng>(&self, rng: &mut R) -> StatsResult<Array2<u32>> {
let mut matrix = Array2::zeros((32, 32));
for i in 0..32 {
let j = rng.random_range(0..32);
matrix[[i, j]] = 1;
}
Ok(matrix)
}
fn apply_owen_scrambling(&self, value: u32, matrix: &Array2<u32>) -> u32 {
let mut result = 0u32;
for i in 0..32 {
let bit = (value >> (31 - i)) & 1;
for j in 0..32 {
if matrix[[i, j]] == 1 && bit == 1 {
result |= 1u32 << (31 - j);
break;
}
}
}
result
}
fn smallest_prime_geq(n: u32) -> u32 {
if n <= 2 {
return 2;
}
let mut candidate = if n.is_multiple_of(2) { n + 1 } else { n };
while !Self::is_prime(candidate) {
candidate += 2;
}
candidate
}
fn is_prime(n: u32) -> bool {
if n < 2 {
return false;
}
if n == 2 {
return true;
}
if n.is_multiple_of(2) {
return false;
}
let sqrt_n = (n as f64).sqrt() as u32;
for i in (3..=sqrt_n).step_by(2) {
if n.is_multiple_of(i) {
return false;
}
}
true
}
fn build_faure_pascal_matrix(base: u64) -> Array2<u64> {
let mut binom = [[0u64; FAURE_DIGITS]; FAURE_DIGITS];
for l in 0..FAURE_DIGITS {
binom[l][0] = 1;
binom[l][l] = 1;
for r in 1..l {
binom[l][r] = binom[l - 1][r - 1] + binom[l - 1][r];
}
}
let mut matrix = Array2::<u64>::zeros((FAURE_DIGITS, FAURE_DIGITS));
for r in 0..FAURE_DIGITS {
for l in r..FAURE_DIGITS {
matrix[[r, l]] = binom[l][r] % base;
}
}
matrix
}
fn assess_quality(&mut self, sequence: &Array2<F>) -> StatsResult<()> {
let n = sequence.nrows();
let d = sequence.ncols();
let mut max_discrepancy = 0.0;
let num_test_points = 50.min(n);
let mut rng = scirs2_core::random::thread_rng();
for _ in 0..num_test_points {
let mut test_point = Array1::zeros(d);
for j in 0..d {
test_point[j] = F::from(rng.random::<f64>()).expect("Operation failed");
}
let mut count = 0;
for i in 0..n {
let mut in_box = true;
for j in 0..d {
if sequence[[i, j]] > test_point[j] {
in_box = false;
break;
}
}
if in_box {
count += 1;
}
}
let volume: F = test_point.iter().fold(F::one(), |acc, &x| acc * x);
let expected = volume.to_f64().expect("Operation failed") * n as f64;
let discrepancy = (count as f64 - expected).abs() / n as f64;
max_discrepancy = max_discrepancy.max(discrepancy);
}
self.state.quality_metrics.star_discrepancy = max_discrepancy;
let m = n.min(QUALITY_SAMPLE_CAP);
let d_wd = d.min(WRAPAROUND_DIM_CAP);
let mut wd_sum = 0.0f64;
for i in 0..m {
for k in 0..m {
let mut prod = 1.0f64;
for j in 0..d_wd {
let xi = sequence[[i, j]].to_f64().expect("Operation failed");
let xk = sequence[[k, j]].to_f64().expect("Operation failed");
let diff = (xi - xk).abs();
prod *= 1.5 - diff * (1.0 - diff);
}
wd_sum += prod;
}
}
let wd_squared = -((4.0f64 / 3.0).powi(d_wd as i32)) + wd_sum / (m as f64 * m as f64);
self.state.quality_metrics.wraparound_discrepancy = wd_squared.max(0.0).sqrt();
let dims = d.min(DIAPHONY_DIM_CAP);
let mut diaphony_sum = 0.0f64;
for j in 0..dims {
for h in 1..=DIAPHONY_FREQ_MAX {
let (re, im) = Self::empirical_fourier_coefficient(sequence, m, &[(j, h)]);
let weight = 1.0 / (h as f64 * h as f64);
diaphony_sum += weight * (re * re + im * im);
}
}
for j1 in 0..dims {
for j2 in (j1 + 1)..dims {
for h1 in 1..=DIAPHONY_FREQ_MAX {
for h2 in -DIAPHONY_FREQ_MAX..=DIAPHONY_FREQ_MAX {
if h2 == 0 {
continue;
}
let (re, im) =
Self::empirical_fourier_coefficient(sequence, m, &[(j1, h1), (j2, h2)]);
let weight = 1.0 / ((h1 * h1) as f64 * (h2 * h2) as f64);
diaphony_sum += weight * (re * re + im * im);
}
}
}
}
self.state.quality_metrics.diaphony = diaphony_sum.sqrt();
self.state.quality_metrics.figure_of_merit = self
.state
.quality_metrics
.star_discrepancy
.max(self.state.quality_metrics.wraparound_discrepancy)
.max(self.state.quality_metrics.diaphony);
Ok(())
}
fn empirical_fourier_coefficient(
sequence: &Array2<F>,
m: usize,
terms: &[(usize, i64)],
) -> (f64, f64) {
let mut re = 0.0f64;
let mut im = 0.0f64;
for k in 0..m {
let mut phase = 0.0f64;
for &(dim, h) in terms {
let x = sequence[[k, dim]].to_f64().expect("Operation failed");
phase += h as f64 * x;
}
let angle = 2.0 * std::f64::consts::PI * phase;
re += angle.cos();
im += angle.sin();
}
(re / m as f64, im / m as f64)
}
pub fn quality_metrics(&self) -> &QualityMetrics {
&self.state.quality_metrics
}
}
#[allow(dead_code)]
pub fn enhanced_sobol<F>(
n: usize,
dimension: usize,
scrambling: bool,
seed: Option<u64>,
) -> StatsResult<Array2<F>>
where
F: Float + Zero + One + Copy + Send + Sync + SimdUnifiedOps + FromPrimitive + std::fmt::Display,
{
let sequence_type = EnhancedSequenceType::SobolAdvanced {
owen_scrambling: scrambling,
digital_shift: true,
nested_scrambling: false,
};
let config = EnhancedQMCConfig {
seed,
..Default::default()
};
let mut generator = EnhancedQMCGenerator::new(sequence_type, dimension, config)?;
generator.generate(n)
}
#[allow(dead_code)]
pub fn enhanced_niederreiter<F>(
n: usize,
dimension: usize,
seed: Option<u64>,
) -> StatsResult<Array2<F>>
where
F: Float + Zero + One + Copy + Send + Sync + SimdUnifiedOps + FromPrimitive + std::fmt::Display,
{
let sequence_type = EnhancedSequenceType::Niederreiter {
base_strategy: BaseSelectionStrategy::OptimizedPrimes,
matrix_optimization: true,
};
let config = EnhancedQMCConfig {
seed,
..Default::default()
};
let mut generator = EnhancedQMCGenerator::new(sequence_type, dimension, config)?;
generator.generate(n)
}
#[allow(dead_code)]
pub fn enhanced_digital_net<F>(
n: usize,
dimension: usize,
t: usize,
seed: Option<u64>,
) -> StatsResult<Array2<F>>
where
F: Float + Zero + One + Copy + Send + Sync + SimdUnifiedOps + FromPrimitive + std::fmt::Display,
{
let net_params = DigitalNetParams {
t,
m: 32,
s: dimension,
base: 2,
};
let sequence_type = EnhancedSequenceType::DigitalNet {
net_params,
construction_method: NetConstructionMethod::Sobol,
};
let config = EnhancedQMCConfig {
seed,
..Default::default()
};
let mut generator = EnhancedQMCGenerator::new(sequence_type, dimension, config)?;
generator.generate(n)
}
#[path = "enhanced_sequences_tests.rs"]
#[cfg(test)]
mod tests;