use super::cert::SosPoly;
use super::lp::{Lp, LpStatus, Rel};
use super::ratpoly::{Exponents, RatPoly};
use rug::Rational;
use std::collections::{BTreeMap, BTreeSet};
pub fn monomial_basis(nvars: usize, max_deg: u32) -> Vec<Exponents> {
let mut out = Vec::new();
if nvars == 0 {
out.push(Vec::new());
return out;
}
let mut cur = vec![0u32; nvars];
fn rec(idx: usize, nvars: usize, remaining: u32, cur: &mut Vec<u32>, out: &mut Vec<Exponents>) {
if idx == nvars {
out.push(cur.clone());
return;
}
for k in 0..=remaining {
cur[idx] = k;
rec(idx + 1, nvars, remaining - k, cur, out);
}
cur[idx] = 0;
}
rec(0, nvars, max_deg, &mut cur, &mut out);
out.sort_by(|a, b| {
let da: u32 = a.iter().sum();
let db: u32 = b.iter().sum();
da.cmp(&db).then_with(|| a.cmp(b))
});
out
}
fn add_exp(a: &[u32], b: &[u32]) -> Exponents {
a.iter().zip(b).map(|(x, y)| x + y).collect()
}
const RATIOS: &[(i32, i32)] = &[
(1, 1),
(1, 2),
(2, 1),
(1, 3),
(3, 1),
(2, 3),
(3, 2),
(1, 4),
(4, 1),
];
const MAX_LP_COLUMNS: usize = 1200;
fn generators(n: usize) -> Vec<Vec<(usize, Rational)>> {
let mut gens: Vec<Vec<(usize, Rational)>> = Vec::new();
for i in 0..n {
gens.push(vec![(i, Rational::from(1))]);
}
let pairs: Vec<(usize, usize)> = (0..n)
.flat_map(|i| ((i + 1)..n).map(move |j| (i, j)))
.collect();
let full = n + pairs.len() * RATIOS.len() * 2;
let ratios: &[(i32, i32)] = if full > MAX_LP_COLUMNS {
&RATIOS[..1]
} else {
RATIOS
};
for &(i, j) in &pairs {
for &(a, b) in ratios {
for sign in [1, -1] {
gens.push(vec![(i, Rational::from(a)), (j, Rational::from(sign * b))]);
}
}
}
gens
}
pub fn dsos_search(p: &RatPoly, basis_deg: u32) -> Option<SosPoly> {
let nvars = p.nvars();
let basis = monomial_basis(nvars, basis_deg);
let n = basis.len();
if n == 0 {
return if p.is_zero() {
Some(SosPoly::default())
} else {
None
};
}
let gens = generators(n);
let mut acc: BTreeMap<Exponents, Vec<(usize, Rational)>> = BTreeMap::new();
for (g_idx, g) in gens.iter().enumerate() {
let mut local: BTreeMap<Exponents, Rational> = BTreeMap::new();
for (u, cu) in g {
for (v, cv) in g {
let e = add_exp(&basis[*u], &basis[*v]);
*local.entry(e).or_insert_with(|| Rational::from(0)) += cu.clone() * cv.clone();
}
}
for (e, c) in local {
if c != 0 {
acc.entry(e).or_default().push((g_idx, c));
}
}
}
let mut all_exps: BTreeSet<Exponents> = acc.keys().cloned().collect();
all_exps.extend(p.terms().keys().cloned());
let mut lp = Lp::new(gens.len());
for exp in &all_exps {
let mut row = vec![Rational::from(0); gens.len()];
if let Some(contribs) = acc.get(exp) {
for (idx, c) in contribs {
row[*idx] += c.clone();
}
}
lp.constrain(row, Rel::Eq, p.coeff(exp));
}
for k in 0..gens.len() {
lp.set_objective(k, Rational::from(1));
}
let x = match lp.solve() {
LpStatus::Optimal(x) => x,
_ => return None,
};
let mut sos = SosPoly::default();
for (g_idx, g) in gens.iter().enumerate() {
let w = x[g_idx].clone();
if w <= 0 {
continue;
}
let mut square = RatPoly::zero(nvars);
for (u, c) in g {
square = square.add(&RatPoly::monomial(nvars, basis[*u].clone(), c.clone()));
}
sos.push(w, square);
}
Some(sos)
}
#[cfg(test)]
mod tests {
use super::*;
fn r(n: i64, d: i64) -> Rational {
Rational::from((n, d))
}
#[test]
fn monomial_basis_univariate() {
let b = monomial_basis(1, 2);
assert_eq!(b, vec![vec![0], vec![1], vec![2]]);
}
#[test]
fn monomial_basis_bivariate_degree1() {
let b = monomial_basis(2, 1);
let mut expect = vec![vec![0, 0], vec![0, 1], vec![1, 0]];
expect.sort();
let mut got = b.clone();
got.sort();
assert_eq!(got, expect);
}
#[test]
fn dsos_finds_perfect_square() {
let mut p = RatPoly::monomial(1, vec![2], Rational::from(1));
p = p.add(&RatPoly::monomial(1, vec![1], Rational::from(2)));
p = p.add(&RatPoly::constant(1, Rational::from(1)));
let sos = dsos_search(&p, 1).expect("DSOS should find a certificate");
assert_eq!(sos.to_poly(1), p);
}
#[test]
fn dsos_finds_diagonal_sum() {
let p = RatPoly::monomial(2, vec![2, 0], Rational::from(1)).add(&RatPoly::monomial(
2,
vec![0, 2],
Rational::from(1),
));
let sos = dsos_search(&p, 1).expect("DSOS should find a certificate");
assert_eq!(sos.to_poly(2), p);
}
#[test]
fn dsos_refuses_unreachable_degree() {
let p = RatPoly::monomial(1, vec![4], Rational::from(1));
assert!(dsos_search(&p, 1).is_none());
}
#[test]
fn dsos_handles_off_diagonal_quadratic() {
let mut p = RatPoly::monomial(2, vec![2, 0], Rational::from(1));
p = p.add(&RatPoly::monomial(2, vec![1, 1], Rational::from(-2)));
p = p.add(&RatPoly::monomial(2, vec![0, 2], Rational::from(2)));
let sos = dsos_search(&p, 1).expect("DSOS should find a certificate");
assert_eq!(sos.to_poly(2), p);
}
#[test]
fn dsos_rational_coefficients() {
let mut p = RatPoly::monomial(1, vec![2], r(1, 4));
p = p.add(&RatPoly::monomial(1, vec![1], r(1, 3)));
p = p.add(&RatPoly::constant(1, r(1, 9)));
let sos = dsos_search(&p, 1).expect("DSOS should find a certificate");
assert_eq!(sos.to_poly(1), p);
}
}