use std::fmt;
use crate::core::math::Scalar;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum RootTerminationReason {
Converged,
MaxIter,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RootResult<F = f64> {
root: F,
value: F,
lower: F,
upper: F,
iterations: u64,
function_evals: u64,
reason: RootTerminationReason,
}
impl<F: Scalar> RootResult<F> {
fn new(
root: F,
value: F,
bracket: (F, F),
work: (u64, u64),
reason: RootTerminationReason,
) -> Self {
let (a, b) = bracket;
let (iterations, function_evals) = work;
let (lower, upper) = if a <= b { (a, b) } else { (b, a) };
Self {
root,
value,
lower,
upper,
iterations,
function_evals,
reason,
}
}
pub fn root(&self) -> F {
self.root
}
pub fn value(&self) -> F {
self.value
}
pub fn bracket(&self) -> (F, F) {
(self.lower, self.upper)
}
pub fn iterations(&self) -> u64 {
self.iterations
}
pub fn function_evals(&self) -> u64 {
self.function_evals
}
pub fn reason(&self) -> RootTerminationReason {
self.reason
}
pub fn converged(&self) -> bool {
self.reason == RootTerminationReason::Converged
}
}
#[derive(Debug, PartialEq)]
#[non_exhaustive]
pub enum BrentRootError<E, F = f64> {
Evaluation(E),
InvalidInterval {
lower: F,
upper: F,
},
NotBracketed {
lower: F,
upper: F,
f_lower: F,
f_upper: F,
},
NonFiniteValue {
x: F,
value: F,
},
}
impl<E, F> fmt::Display for BrentRootError<E, F>
where
E: fmt::Display,
F: fmt::Debug,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Evaluation(error) => {
write!(formatter, "root function failed: {error}")
}
Self::InvalidInterval { lower, upper } => write!(
formatter,
"root interval must be finite and ordered, got [{lower:?}, {upper:?}]"
),
Self::NotBracketed {
lower,
upper,
f_lower,
f_upper,
} => write!(
formatter,
"root is not bracketed on [{lower:?}, {upper:?}]: endpoint values are {f_lower:?} and {f_upper:?}"
),
Self::NonFiniteValue { x, value } => write!(
formatter,
"root function returned non-finite value {value:?} at {x:?}"
),
}
}
}
impl<E, F> std::error::Error for BrentRootError<E, F>
where
E: std::error::Error + 'static,
F: fmt::Debug + 'static,
{
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Evaluation(error) => Some(error),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct BrentRoot<F = f64> {
lower: F,
upper: F,
tol_rel: F,
tol_abs: F,
max_iter: u64,
}
impl<F: Scalar> BrentRoot<F> {
pub fn new(lower: F, upper: F) -> Self {
Self {
lower,
upper,
tol_rel: F::from_f64(4.0).unwrap() * F::epsilon(),
tol_abs: F::from_f64(1e-12).unwrap(),
max_iter: 100,
}
}
pub fn with_tol(mut self, tol_rel: F, tol_abs: F) -> Self {
let min_tol_rel = F::from_f64(4.0).unwrap() * F::epsilon();
assert!(
tol_rel.is_finite() && tol_rel >= min_tol_rel,
"BrentRoot relative tolerance must be finite and at least four times machine epsilon"
);
assert!(
tol_abs.is_finite() && tol_abs > F::zero(),
"BrentRoot absolute tolerance must be finite and positive"
);
self.tol_rel = tol_rel;
self.tol_abs = tol_abs;
self
}
pub fn with_max_iter(mut self, max_iter: u64) -> Self {
self.max_iter = max_iter;
self
}
pub fn solve<Function, E>(
&self,
mut function: Function,
) -> Result<RootResult<F>, BrentRootError<E, F>>
where
Function: FnMut(F) -> Result<F, E>,
{
if !self.lower.is_finite()
|| !self.upper.is_finite()
|| self.lower >= self.upper
|| !(self.upper - self.lower).is_finite()
{
return Err(BrentRootError::InvalidInterval {
lower: self.lower,
upper: self.upper,
});
}
let mut function_evals = 1;
let mut a = self.lower;
let mut fa = function(a).map_err(BrentRootError::Evaluation)?;
if !fa.is_finite() {
return Err(BrentRootError::NonFiniteValue { x: a, value: fa });
}
if fa == F::zero() {
return Ok(RootResult::new(
a,
fa,
(self.lower, self.upper),
(0, function_evals),
RootTerminationReason::Converged,
));
}
function_evals += 1;
let mut b = self.upper;
let mut fb = function(b).map_err(BrentRootError::Evaluation)?;
if !fb.is_finite() {
return Err(BrentRootError::NonFiniteValue { x: b, value: fb });
}
if fb == F::zero() {
return Ok(RootResult::new(
b,
fb,
(self.lower, self.upper),
(0, function_evals),
RootTerminationReason::Converged,
));
}
if same_nonzero_sign(fa, fb) {
return Err(BrentRootError::NotBracketed {
lower: self.lower,
upper: self.upper,
f_lower: fa,
f_upper: fb,
});
}
let mut c = b;
let mut fc = fb;
let mut d = b - a;
let mut e = d;
let half = F::from_f64(0.5).unwrap();
let two = F::from_f64(2.0).unwrap();
let three = F::from_f64(3.0).unwrap();
for iteration in 0..self.max_iter {
if same_nonzero_sign(fb, fc) {
c = a;
fc = fa;
d = b - a;
e = d;
}
if fc.abs() < fb.abs() {
(a, b, c) = (b, c, b);
(fa, fb, fc) = (fb, fc, fb);
}
let tolerance = half * (self.tol_abs + self.tol_rel * b.abs());
let midpoint = half * (c - b);
if midpoint.abs() <= tolerance || fb == F::zero() {
return Ok(RootResult::new(
b,
fb,
(b, c),
(iteration, function_evals),
RootTerminationReason::Converged,
));
}
if e.abs() >= tolerance && fa.abs() > fb.abs() {
let s = fb / fa;
let (mut p, mut q) = if a == c {
(two * midpoint * s, F::one() - s)
} else {
let q = fa / fc;
let r = fb / fc;
(
s * (two * midpoint * q * (q - r)
- (b - a) * (r - F::one())),
(q - F::one()) * (r - F::one()) * (s - F::one()),
)
};
if p > F::zero() {
q = -q;
}
p = p.abs();
let interpolation_bound =
three * midpoint * q - (tolerance * q).abs();
let history_bound = (e * q).abs();
if two * p < interpolation_bound.min(history_bound) {
e = d;
d = p / q;
} else {
d = midpoint;
e = d;
}
} else {
d = midpoint;
e = d;
}
a = b;
fa = fb;
b = if d.abs() > tolerance {
b + d
} else if midpoint > F::zero() {
b + tolerance
} else {
b - tolerance
};
function_evals += 1;
fb = function(b).map_err(BrentRootError::Evaluation)?;
if !fb.is_finite() {
return Err(BrentRootError::NonFiniteValue { x: b, value: fb });
}
if fb == F::zero() {
return Ok(RootResult::new(
b,
fb,
(b, b),
(iteration + 1, function_evals),
RootTerminationReason::Converged,
));
}
}
if same_nonzero_sign(fb, fc) {
c = a;
fc = fa;
}
if fc.abs() < fb.abs() {
std::mem::swap(&mut b, &mut c);
fb = fc;
}
Ok(RootResult::new(
b,
fb,
(b, c),
(self.max_iter, function_evals),
RootTerminationReason::MaxIter,
))
}
}
fn same_nonzero_sign<F: Scalar>(a: F, b: F) -> bool {
(a > F::zero() && b > F::zero()) || (a < F::zero() && b < F::zero())
}