use core::fmt;
use core::ops::Deref;
use crate::Interval;
pub const POLYNOMIAL_ROUNDING: f64 = 256.0 * f64::EPSILON;
const MAX_ITERATIONS: usize = 1100;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Root {
pub value: f64,
pub multiplicity: usize,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Roots {
len: usize,
items: [Root; 4],
}
impl Roots {
const EMPTY: Roots = Roots {
len: 0,
items: [Root {
value: 0.0,
multiplicity: 0,
}; 4],
};
fn push(&mut self, root: Root) {
if self.len < self.items.len() {
self.items[self.len] = root;
self.len += 1;
}
}
pub fn as_slice(&self) -> &[Root] {
&self.items[..self.len]
}
pub fn total_multiplicity(&self) -> usize {
self.as_slice().iter().map(|r| r.multiplicity).sum()
}
}
impl Default for Roots {
fn default() -> Self {
Roots::EMPTY
}
}
impl Deref for Roots {
type Target = [Root];
fn deref(&self) -> &[Root] {
self.as_slice()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RootError {
NonFinite,
Zero,
NoSignChange,
}
impl fmt::Display for RootError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
RootError::NonFinite => "root finding on a non-finite input",
RootError::Zero => "every coefficient is zero: every point is a root",
RootError::NoSignChange => "the function has the same sign at both bracket ends",
})
}
}
impl std::error::Error for RootError {}
pub fn quadratic(a: f64, b: f64, c: f64) -> Result<Roots, RootError> {
real_roots(&[c, b, a])
}
pub fn cubic(a: f64, b: f64, c: f64, d: f64) -> Result<Roots, RootError> {
real_roots(&[d, c, b, a])
}
pub fn quartic(a: f64, b: f64, c: f64, d: f64, e: f64) -> Result<Roots, RootError> {
real_roots(&[e, d, c, b, a])
}
pub fn newton_in_interval(
f: impl Fn(f64) -> f64,
df: impl Fn(f64) -> f64,
bracket: Interval,
tol: f64,
) -> Result<f64, RootError> {
if !(bracket.is_bounded() && tol.is_finite() && tol >= 0.0) {
return Err(RootError::NonFinite);
}
let (lo, hi) = (bracket.lo(), bracket.hi());
let (f_lo, f_hi) = (f(lo), f(hi));
if !(f_lo.is_finite() && f_hi.is_finite()) {
return Err(RootError::NonFinite);
}
if f_lo == 0.0 {
return Ok(lo);
}
if f_hi == 0.0 {
return Ok(hi);
}
if (f_lo < 0.0) == (f_hi < 0.0) {
return Err(RootError::NoSignChange);
}
bracketed_newton(&f, &df, lo, hi, f_lo < 0.0, tol)
}
fn bracketed_newton(
f: &dyn Fn(f64) -> f64,
df: &dyn Fn(f64) -> f64,
mut lo: f64,
mut hi: f64,
lo_negative: bool,
tol: f64,
) -> Result<f64, RootError> {
let mut x = 0.5 * (lo + hi);
for _ in 0..MAX_ITERATIONS {
let fx = f(x);
if !fx.is_finite() {
return Err(RootError::NonFinite);
}
if fx == 0.0 {
return Ok(x);
}
if (fx < 0.0) == lo_negative {
lo = x;
} else {
hi = x;
}
let width = hi - lo;
if width <= tol.max(f64::EPSILON * lo.abs().max(hi.abs())) {
return Ok(x);
}
let d = df(x);
let newton = x - fx / d;
let next = if d != 0.0 && newton > lo && newton < hi {
newton
} else {
0.5 * (lo + hi)
};
if next == x {
return Ok(x);
}
x = next;
}
Ok(x)
}
fn eval(c: &[f64], x: f64) -> f64 {
c.iter().rev().fold(0.0, |acc, &k| acc * x + k)
}
fn eval_abs(c: &[f64], x: f64) -> f64 {
let x = x.abs();
c.iter().rev().fold(0.0, |acc, &k| acc * x + k.abs())
}
fn cauchy_bound(c: &[f64]) -> Result<f64, RootError> {
let lead = c[c.len() - 1];
let ratio = c[..c.len() - 1]
.iter()
.map(|k| (k / lead).abs())
.fold(0.0, f64::max);
let bound = 1.0 + ratio;
if bound.is_finite() && eval_abs(c, bound).is_finite() {
Ok(bound)
} else {
Err(RootError::NonFinite)
}
}
fn real_roots(c: &[f64]) -> Result<Roots, RootError> {
if c.iter().any(|k| !k.is_finite()) {
return Err(RootError::NonFinite);
}
let degree = match c.iter().rposition(|&k| k != 0.0) {
None => return Err(RootError::Zero),
Some(0) => return Ok(Roots::EMPTY),
Some(n) => n,
};
let c = &c[..=degree];
let mut out = Roots::EMPTY;
if degree == 1 {
out.push(Root {
value: -c[0] / c[1],
multiplicity: 1,
});
return Ok(out);
}
let mut derivative = [0.0; 4];
for (k, d) in derivative.iter_mut().enumerate().take(degree) {
*d = c[k + 1] * (k + 1) as f64;
}
let critical = real_roots(&derivative[..degree])?;
let bound = cauchy_bound(c)?;
let lead_negative = c[degree] < 0.0;
let mut lo = -bound;
let mut lo_negative = lead_negative != (degree % 2 == 1);
let mut skip = false;
let f = |x| eval(c, x);
let df = |x| eval(&derivative[..degree], x);
for cp in critical.iter() {
let x = cp.value;
if x <= -bound || x >= bound {
continue;
}
let fx = f(x);
if fx.abs() <= POLYNOMIAL_ROUNDING * eval_abs(c, x) {
out.push(Root {
value: x,
multiplicity: cp.multiplicity + 1,
});
lo = x;
skip = true;
continue;
}
let negative = fx < 0.0;
if !skip && negative != lo_negative {
out.push(Root {
value: bracketed_newton(&f, &df, lo, x, lo_negative, 0.0)?,
multiplicity: 1,
});
}
lo = x;
lo_negative = negative;
skip = false;
}
if !skip && lead_negative != lo_negative {
out.push(Root {
value: bracketed_newton(&f, &df, lo, bound, lo_negative, 0.0)?,
multiplicity: 1,
});
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn leading_zeros_lower_the_degree() {
let r = quartic(0.0, 0.0, 1.0, 0.0, -4.0).unwrap();
assert_eq!(r.len(), 2);
assert_eq!((r[0].value, r[1].value), (-2.0, 2.0));
let linear = cubic(0.0, 0.0, 2.0, -1.0).unwrap();
assert_eq!(
linear.as_slice(),
&[Root {
value: 0.5,
multiplicity: 1
}]
);
assert!(quadratic(0.0, 0.0, 3.0).unwrap().is_empty());
}
#[test]
fn degenerate_inputs_are_errors() {
assert_eq!(quadratic(0.0, 0.0, 0.0), Err(RootError::Zero));
assert_eq!(quadratic(f64::NAN, 1.0, 0.0), Err(RootError::NonFinite));
assert_eq!(
quartic(1e-300, 0.0, 0.0, 0.0, 1e300),
Err(RootError::NonFinite)
);
let f = |x: f64| x * x + 1.0;
assert_eq!(
newton_in_interval(f, |x| 2.0 * x, Interval::UNIT, 0.0),
Err(RootError::NoSignChange)
);
assert_eq!(
newton_in_interval(f, |x| 2.0 * x, Interval::REAL, 0.0),
Err(RootError::NonFinite)
);
}
#[test]
fn a_triple_root_is_found_once() {
let r = cubic(1.0, -6.0, 12.0, -8.0).unwrap();
assert_eq!(r.len(), 1);
assert_eq!(r[0].multiplicity, 3);
assert!((r[0].value - 2.0).abs() < 1e-14);
assert_eq!(r.total_multiplicity(), 3);
}
#[test]
fn newton_returns_an_end_that_is_exactly_a_root() {
let f = |x: f64| x - 1.0;
let bracket = Interval::new(1.0, 3.0).unwrap();
assert_eq!(newton_in_interval(f, |_| 1.0, bracket, 0.0), Ok(1.0));
let bracket = Interval::new(-1.0, 1.0).unwrap();
assert_eq!(newton_in_interval(f, |_| 1.0, bracket, 0.0), Ok(1.0));
}
}