use super::{Monomial, MonomialOrder, Polynomial, Term, Var};
#[allow(unused_imports)]
use crate::prelude::*;
use num_rational::BigRational;
use num_traits::{One, Zero};
#[derive(Debug, Clone)]
pub struct SparseConfig {
pub sparsity_threshold: f64,
pub enable_fast_mul: bool,
pub max_sparse_terms: usize,
}
impl Default for SparseConfig {
fn default() -> Self {
Self {
sparsity_threshold: 0.1,
enable_fast_mul: true,
max_sparse_terms: 10000,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SparseStats {
pub zero_terms_skipped: u64,
pub memory_saved: u64,
pub fast_muls: u64,
}
pub struct SparseOps {
config: SparseConfig,
stats: SparseStats,
}
impl SparseOps {
pub fn new(config: SparseConfig) -> Self {
Self {
config,
stats: SparseStats::default(),
}
}
pub fn default_config() -> Self {
Self::new(SparseConfig::default())
}
pub fn is_sparse(&self, p: &Polynomial) -> bool {
if p.num_terms() > self.config.max_sparse_terms {
return false;
}
let num_vars = p.vars().len();
let max_degree = p.total_degree() as usize;
if num_vars == 0 {
return false;
}
let approx_dense_size = if num_vars <= 3 {
max_degree.pow(num_vars as u32)
} else {
max_degree * num_vars * 100
};
let sparsity = p.num_terms() as f64 / approx_dense_size as f64;
sparsity < self.config.sparsity_threshold
}
pub fn sparse_mul(&mut self, p: &Polynomial, q: &Polynomial) -> Polynomial {
if !self.config.enable_fast_mul {
return p * q; }
self.stats.fast_muls += 1;
let mut term_map: FxHashMap<Monomial, BigRational> = FxHashMap::default();
for term_p in p.terms() {
for term_q in q.terms() {
let mono = term_p.monomial.mul(&term_q.monomial);
let coeff = &term_p.coeff * &term_q.coeff;
if !coeff.is_zero() {
term_map
.entry(mono)
.and_modify(|c| *c = c.clone() + &coeff)
.or_insert(coeff);
} else {
self.stats.zero_terms_skipped += 1;
}
}
}
let mut terms: Vec<Term> = term_map
.iter()
.filter_map(|(mono, coeff)| {
if !coeff.is_zero() {
Some(Term::new(coeff.clone(), mono.clone()))
} else {
self.stats.zero_terms_skipped += 1;
None
}
})
.collect();
terms.sort_by(|a, b| MonomialOrder::GRevLex.compare(&a.monomial, &b.monomial));
Polynomial::from_terms(terms, MonomialOrder::GRevLex)
}
pub fn sparse_add(&mut self, p: &Polynomial, q: &Polynomial) -> Polynomial {
let mut term_map: FxHashMap<Monomial, BigRational> = FxHashMap::default();
for term in p.terms() {
term_map.insert(term.monomial.clone(), term.coeff.clone());
}
for term in q.terms() {
term_map
.entry(term.monomial.clone())
.and_modify(|c| *c = c.clone() + &term.coeff)
.or_insert(term.coeff.clone());
}
let terms: Vec<Term> = term_map
.iter()
.filter_map(|(mono, coeff)| {
if !coeff.is_zero() {
Some(Term::new(coeff.clone(), mono.clone()))
} else {
self.stats.zero_terms_skipped += 1;
None
}
})
.collect();
Polynomial::from_terms(terms, MonomialOrder::GRevLex)
}
pub fn sparse_eval(
&mut self,
p: &Polynomial,
point: &FxHashMap<Var, BigRational>,
) -> BigRational {
let mut result = BigRational::zero();
for term in p.terms() {
let mut mono_val = BigRational::one();
for vp in term.monomial.vars() {
if let Some(val) = point.get(&vp.var) {
let powered = self.power_rational(val, vp.power);
mono_val *= powered;
} else {
self.stats.zero_terms_skipped += 1;
mono_val = BigRational::zero();
break;
}
}
result += &term.coeff * &mono_val;
}
result
}
fn power_rational(&self, base: &BigRational, exp: u32) -> BigRational {
if exp == 0 {
BigRational::one()
} else if exp == 1 {
base.clone()
} else {
let mut result = BigRational::one();
let mut b = base.clone();
let mut e = exp;
while e > 0 {
if e % 2 == 1 {
result *= &b;
}
b = &b * &b;
e /= 2;
}
result
}
}
pub fn estimate_memory(&self, p: &Polynomial) -> usize {
p.num_terms() * 100
}
pub fn estimate_savings(&mut self, p: &Polynomial) -> usize {
let num_vars = p.vars().len();
let max_degree = p.total_degree() as usize;
let dense_terms = if num_vars <= 3 {
max_degree.pow(num_vars as u32)
} else {
max_degree * num_vars * 100
};
let sparse_memory = self.estimate_memory(p);
let dense_memory = dense_terms * 100;
let savings = dense_memory.saturating_sub(sparse_memory);
self.stats.memory_saved += savings as u64;
savings
}
pub fn stats(&self) -> &SparseStats {
&self.stats
}
pub fn reset_stats(&mut self) {
self.stats = SparseStats::default();
}
}
#[cfg(test)]
mod tests {
use super::*;
use num_bigint::BigInt;
fn rat(n: i64) -> BigRational {
BigRational::from_integer(BigInt::from(n))
}
#[test]
fn test_sparse_ops_creation() {
let ops = SparseOps::default_config();
assert_eq!(ops.stats().fast_muls, 0);
}
#[test]
fn test_is_sparse() {
let ops = SparseOps::default_config();
let sparse = Polynomial::from_coeffs_int(&[(1, &[(0, 5)]), (1, &[(1, 5)])]);
assert!(ops.is_sparse(&sparse));
let constant = Polynomial::constant(BigRational::from_integer(BigInt::from(5)));
assert!(!ops.is_sparse(&constant));
}
#[test]
fn test_sparse_mul() {
let mut ops = SparseOps::default_config();
let p = Polynomial::from_var(0);
let q = Polynomial::from_var(1);
let result = ops.sparse_mul(&p, &q);
assert_eq!(result.total_degree(), 2);
assert_eq!(ops.stats().fast_muls, 1);
}
#[test]
fn test_sparse_add() {
let mut ops = SparseOps::default_config();
let p = Polynomial::from_var(0);
let q = Polynomial::from_var(1);
let result = ops.sparse_add(&p, &q);
assert_eq!(result.num_terms(), 2);
}
#[test]
fn test_sparse_eval() {
let mut ops = SparseOps::default_config();
let p = Polynomial::from_coeffs_int(&[(2, &[(0, 1)]), (3, &[(1, 1)])]);
let mut point = FxHashMap::default();
point.insert(0, rat(5)); point.insert(1, rat(2));
let result = ops.sparse_eval(&p, &point);
assert_eq!(result, rat(16));
}
#[test]
fn test_power_rational() {
let ops = SparseOps::default_config();
assert_eq!(ops.power_rational(&rat(2), 0), rat(1));
assert_eq!(ops.power_rational(&rat(2), 1), rat(2));
assert_eq!(ops.power_rational(&rat(2), 3), rat(8));
}
#[test]
fn test_estimate_memory() {
let ops = SparseOps::default_config();
let p = Polynomial::from_coeffs_int(&[(1, &[(0, 1)]), (1, &[(1, 1)])]);
let memory = ops.estimate_memory(&p);
assert!(memory > 0);
}
#[test]
fn test_estimate_savings() {
let mut ops = SparseOps::default_config();
let p = Polynomial::from_coeffs_int(&[
(1, &[(0, 10)]), (1, &[]), ]);
let savings = ops.estimate_savings(&p);
assert!(savings > 0);
}
}