use crate::{Cubic, InputError, TerminationCondition, different_signs};
#[derive(Clone, Debug)]
pub struct Poly {
coeffs: Vec<f64>,
}
impl<'a> std::ops::Mul<&'a Poly> for &'a Poly {
type Output = Poly;
fn mul(self, rhs: &Poly) -> Poly {
let mut coeffs = vec![0.0; (self.coeffs.len() + rhs.coeffs.len()).saturating_sub(1)];
for (i, c) in self.coeffs.iter().enumerate() {
for (j, d) in rhs.coeffs.iter().enumerate() {
coeffs[i + j] += c * d;
}
}
Poly { coeffs }
}
}
impl std::ops::Mul<&Poly> for Poly {
type Output = Poly;
fn mul(self, rhs: &Poly) -> Poly {
(&self) * rhs
}
}
impl Poly {
pub fn new(coeffs: impl IntoIterator<Item = f64>) -> Self {
Poly {
coeffs: coeffs.into_iter().collect(),
}
}
fn is_finite(&self) -> bool {
self.coeffs.iter().all(|c| c.is_finite())
}
pub fn deriv(&self) -> Poly {
let mut coeffs = Vec::with_capacity(self.coeffs.len() - 1);
for (i, c) in self.coeffs.iter().enumerate().skip(1) {
coeffs.push(c * (i as f64));
}
Poly { coeffs }
}
pub fn eval(&self, x: f64) -> f64 {
let mut ret = 0.0;
let mut x_pow = 1.0;
for &c in &self.coeffs {
ret += c * x_pow;
x_pow *= x;
}
ret
}
pub fn degree(&self) -> usize {
self.coeffs.len().saturating_sub(1)
}
pub fn to_cubic(&self) -> Option<Cubic> {
if self.degree() <= 3 {
Some(Cubic {
c0: self.coeffs.first().copied().unwrap_or(0.0),
c1: self.coeffs.get(1).copied().unwrap_or(0.0),
c2: self.coeffs.get(2).copied().unwrap_or(0.0),
c3: self.coeffs.get(3).copied().unwrap_or(0.0),
})
} else {
None
}
}
fn one_root<Term: TerminationCondition>(
&self,
deriv: &Poly,
mut lower: f64,
mut upper: f64,
val_lower: f64,
val_upper: f64,
term: Term,
) -> f64 {
if !val_lower.is_finite() || !val_upper.is_finite() || !deriv.is_finite() {
return f64::NAN;
}
debug_assert!(different_signs(val_lower, val_upper));
let mut x = lower + (upper - lower) / 2.0;
let mut val_x = self.eval(x);
let mut step = (upper - lower) / 2.0;
while x.is_finite() && !term.stop(step, val_x) {
let root_in_first_half = different_signs(val_lower, val_x);
if root_in_first_half {
upper = x;
} else {
lower = x;
}
let deriv_x = self.deriv().eval(x);
debug_assert!(deriv_x.is_finite());
debug_assert!(val_x.is_finite());
step = -val_x / deriv_x;
let mut new_x = x + step;
if new_x <= lower || new_x >= upper {
new_x = lower + (upper - lower) / 2.0;
if new_x == upper || new_x == lower {
return new_x;
}
}
step = new_x - x;
x = new_x;
val_x = self.eval(x);
}
x
}
pub fn roots_in(&self, lower: f64, upper: f64, x_error: f64) -> Vec<f64> {
let mut ret = Vec::new();
if let Some(c) = self.to_cubic() {
ret.extend(c.all_roots(lower, upper, x_error));
return ret;
}
let deriv = self.deriv();
let mut possible_endpoints = deriv.roots_in(lower, upper, x_error);
possible_endpoints.push(upper);
let mut last = lower;
let mut last_val = self.eval(last);
for x in possible_endpoints {
if x > last && x <= upper {
let val = self.eval(x);
if different_signs(last_val, val) {
ret.push(self.one_root(&deriv, last, x, last_val, val, InputError(x_error)));
}
last = x;
last_val = val;
}
}
ret
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn smoke() {
let x_minus_1 = Poly::new([-1.0, 1.0]);
let x_minus_2 = Poly::new([-2.0, 1.0]);
let x_minus_3 = Poly::new([-3.0, 1.0]);
let x_minus_4 = Poly::new([-4.0, 1.0]);
let p = &x_minus_1 * &x_minus_2 * &x_minus_3 * &x_minus_4;
let roots = p.roots_in(0.0, 5.0, 1e-6);
assert_eq!(roots.len(), 4);
assert!((roots[0] - 1.0).abs() <= 1e-6);
assert!((roots[1] - 2.0).abs() <= 1e-6);
assert!((roots[2] - 3.0).abs() <= 1e-6);
assert!((roots[3] - 4.0).abs() <= 1e-6);
}
}