use axiolid_guarantees::Sign;
use crate::arith::{sign_product, Arith};
use crate::certify::ExactError;
pub const MAX_DEPTH: usize = 6;
#[derive(Debug, Clone, PartialEq)]
pub struct Nested<T> {
coeffs: Vec<T>,
}
impl<T> Nested<T> {
#[must_use]
pub fn level(&self) -> usize {
self.coeffs.len().trailing_zeros() as usize
}
#[must_use]
pub fn coeffs(&self) -> &[T] {
&self.coeffs
}
}
#[derive(Debug, Clone, Default)]
pub struct Tower<T> {
radicands: Vec<Nested<T>>,
}
impl<T: Arith> Tower<T> {
#[must_use]
pub fn new() -> Self {
Self {
radicands: Vec::new(),
}
}
#[must_use]
pub fn depth(&self) -> usize {
self.radicands.len()
}
#[must_use]
pub fn value(&self, value: T) -> Nested<T> {
Nested {
coeffs: vec![value],
}
}
#[must_use]
pub fn from_f64(&self, value: f64) -> Nested<T> {
self.value(T::from_f64(value))
}
pub fn sqrt(&mut self, radicand: &Nested<T>) -> Result<Nested<T>, ExactError> {
let level = self.depth();
if level >= MAX_DEPTH {
return Err(ExactError::TooDeep);
}
self.own(radicand);
self.radicands.push(self.lift(radicand, level));
let mut coeffs = vec![T::from_f64(0.0); 1 << (level + 1)];
coeffs[1 << level] = T::from_f64(1.0);
Ok(Nested { coeffs })
}
#[must_use]
pub fn add(&self, x: &Nested<T>, y: &Nested<T>) -> Nested<T> {
self.zip(x, y, T::add)
}
#[must_use]
pub fn sub(&self, x: &Nested<T>, y: &Nested<T>) -> Nested<T> {
self.zip(x, y, T::sub)
}
#[must_use]
pub fn neg(&self, x: &Nested<T>) -> Nested<T> {
self.own(x);
Nested {
coeffs: x.coeffs.iter().map(T::neg).collect(),
}
}
#[must_use]
pub fn mul(&self, x: &Nested<T>, y: &Nested<T>) -> Nested<T> {
self.own(x);
self.own(y);
let level = x.level().max(y.level());
let (x, y) = (self.lift(x, level), self.lift(y, level));
Nested {
coeffs: self.mul_at(level, &x.coeffs, &y.coeffs),
}
}
#[must_use]
pub fn sign(&self, x: &Nested<T>) -> Option<Sign> {
self.own(x);
if let Some(sign) = self
.numeric(x.level(), &x.coeffs)
.and_then(|value| value.sign())
{
return Some(sign);
}
self.sign_at(x.level(), &x.coeffs)
}
fn numeric(&self, level: usize, x: &[T]) -> Option<T> {
if level == 0 {
return Some(x[0].clone());
}
let half = 1 << (level - 1);
let (a, b) = x.split_at(half);
let root = self
.numeric(level - 1, &self.radicands[level - 1].coeffs)?
.sqrt_enclosure()?;
let a = self.numeric(level - 1, a)?;
let b = self.numeric(level - 1, b)?;
Some(a.add(&b.mul(&root)))
}
#[must_use]
pub fn cmp(&self, x: &Nested<T>, y: &Nested<T>) -> Option<Sign> {
self.sign(&self.sub(x, y))
}
fn own(&self, x: &Nested<T>) {
assert!(
x.level() <= self.depth(),
"a Nested value is only meaningful in the tower that made it"
);
}
fn lift(&self, x: &Nested<T>, level: usize) -> Nested<T> {
let mut coeffs = x.coeffs.clone();
coeffs.resize(1 << level, T::from_f64(0.0));
Nested { coeffs }
}
fn zip(&self, x: &Nested<T>, y: &Nested<T>, op: impl Fn(&T, &T) -> T) -> Nested<T> {
self.own(x);
self.own(y);
let level = x.level().max(y.level());
let (x, y) = (self.lift(x, level), self.lift(y, level));
Nested {
coeffs: x
.coeffs
.iter()
.zip(&y.coeffs)
.map(|(a, b)| op(a, b))
.collect(),
}
}
fn mul_at(&self, level: usize, x: &[T], y: &[T]) -> Vec<T> {
if level == 0 {
return vec![x[0].mul(&y[0])];
}
let half = 1 << (level - 1);
let (a, b) = x.split_at(half);
let (c, d) = y.split_at(half);
let radicand = &self.radicands[level - 1].coeffs;
let ac = self.mul_at(level - 1, a, c);
let bd = self.mul_at(level - 1, b, d);
let bdr = self.mul_at(level - 1, &bd, radicand);
let ad = self.mul_at(level - 1, a, d);
let bc = self.mul_at(level - 1, b, c);
let mut out: Vec<T> = ac.iter().zip(&bdr).map(|(p, q)| p.add(q)).collect();
out.extend(ad.iter().zip(&bc).map(|(p, q)| p.add(q)));
out
}
fn sign_at(&self, level: usize, x: &[T]) -> Option<Sign> {
if level == 0 {
return x[0].sign();
}
let half = 1 << (level - 1);
let (a, b) = x.split_at(half);
let radicand = &self.radicands[level - 1].coeffs;
let sr = self.sign_at(level - 1, radicand)?;
if sr == Sign::Negative {
return None;
}
let sa = self.sign_at(level - 1, a)?;
let sb = if sr == Sign::Zero {
Sign::Zero
} else {
self.sign_at(level - 1, b)?
};
if sb == Sign::Zero {
return Some(sa);
}
if sa == Sign::Zero || sa == sb {
return Some(sb);
}
let a2 = self.mul_at(level - 1, a, a);
let b2 = self.mul_at(level - 1, b, b);
let b2r = self.mul_at(level - 1, &b2, radicand);
let diff: Vec<T> = a2.iter().zip(&b2r).map(|(p, q)| p.sub(q)).collect();
let dominance = self.sign_at(level - 1, &diff)?;
Some(sign_product(sa, dominance))
}
}