use thiserror::Error;
#[derive(Debug, Error, PartialEq)]
pub enum RootError {
#[error("interval [{a}, {b}] does not bracket a root: f(a)={fa}, f(b)={fb}")]
NoBracket { a: f64, b: f64, fa: f64, fb: f64 },
#[error("did not converge in {max_iter} iterations (last residual {last_residual})")]
NoConvergence { max_iter: usize, last_residual: f64 },
#[error("f returned NaN at x = {x}")]
NanEvaluation { x: f64 },
}
pub fn brent<F>(mut f: F, a: f64, b: f64, xtol: f64, max_iter: usize) -> Result<f64, RootError>
where
F: FnMut(f64) -> f64,
{
let mut fa = f(a);
let mut fb = f(b);
if fa.is_nan() {
return Err(RootError::NanEvaluation { x: a });
}
if fb.is_nan() {
return Err(RootError::NanEvaluation { x: b });
}
if fa * fb > 0.0 {
return Err(RootError::NoBracket { a, b, fa, fb });
}
let (mut a, mut b) = (a, b);
if fa.abs() < fb.abs() {
std::mem::swap(&mut a, &mut b);
std::mem::swap(&mut fa, &mut fb);
}
let mut c = a;
let mut fc = fa;
let mut d = c;
let mut mflag = true;
for iter in 0..max_iter {
let tol = 2.0 * f64::EPSILON * b.abs() + 0.5 * xtol;
let mid = 0.5 * (a - b);
if fb == 0.0 || mid.abs() < tol {
return Ok(b);
}
let s: f64 = if fa != fc && fb != fc {
let denom_a = (fa - fb) * (fa - fc);
let denom_b = (fb - fa) * (fb - fc);
let denom_c = (fc - fa) * (fc - fb);
a * fb * fc / denom_a + b * fa * fc / denom_b + c * fa * fb / denom_c
} else {
b - fb * (b - a) / (fb - fa)
};
let cond1 = {
let lo = (3.0 * a + b) / 4.0;
let (low, high) = if lo < b { (lo, b) } else { (b, lo) };
s < low || s > high
};
let cond2 = mflag && (s - b).abs() >= (b - c).abs() / 2.0;
let cond3 = !mflag && (s - b).abs() >= (c - d).abs() / 2.0;
let cond4 = mflag && (b - c).abs() < tol;
let cond5 = !mflag && (c - d).abs() < tol;
let s = if cond1 || cond2 || cond3 || cond4 || cond5 {
mflag = true;
0.5 * (a + b)
} else {
mflag = false;
s
};
let fs = f(s);
if fs.is_nan() {
return Err(RootError::NanEvaluation { x: s });
}
d = c;
c = b;
fc = fb;
if fa * fs < 0.0 {
b = s;
fb = fs;
} else {
a = s;
fa = fs;
}
if fa.abs() < fb.abs() {
std::mem::swap(&mut a, &mut b);
std::mem::swap(&mut fa, &mut fb);
}
if iter + 1 == max_iter {
return Err(RootError::NoConvergence {
max_iter,
last_residual: fb,
});
}
}
unreachable!("loop exits via convergence return or NoConvergence")
}
pub fn illinois<F>(mut f: F, a: f64, b: f64, xtol: f64, max_iter: usize) -> Result<f64, RootError>
where
F: FnMut(f64) -> f64,
{
let mut fa = f(a);
let mut fb = f(b);
if fa.is_nan() {
return Err(RootError::NanEvaluation { x: a });
}
if fb.is_nan() {
return Err(RootError::NanEvaluation { x: b });
}
if fa * fb > 0.0 {
return Err(RootError::NoBracket { a, b, fa, fb });
}
let (mut a, mut b) = (a, b);
let mut side: i32 = 0;
for iter in 0..max_iter {
let denom = fb - fa;
let c = if denom == 0.0 {
0.5 * (a + b)
} else {
(a * fb - b * fa) / denom
};
let fc = f(c);
if fc.is_nan() {
return Err(RootError::NanEvaluation { x: c });
}
if fc == 0.0 || (b - a).abs() < xtol {
return Ok(c);
}
if fc * fb < 0.0 {
a = c;
fa = fc;
if side == -1 {
fb *= 0.5; }
side = -1;
} else {
b = c;
fb = fc;
if side == 1 {
fa *= 0.5;
}
side = 1;
}
if iter + 1 == max_iter {
return Err(RootError::NoConvergence {
max_iter,
last_residual: fc,
});
}
}
unreachable!("loop exits via convergence return or NoConvergence")
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_close(actual: f64, expected: f64, tol: f64) {
assert!(
(actual - expected).abs() < tol,
"expected {expected}, got {actual}"
);
}
fn cos_minus_x(x: f64) -> f64 {
x.cos() - x
}
fn wallis(x: f64) -> f64 {
x.powi(3) - 2.0 * x - 5.0
}
#[test]
fn brent_finds_dottie_number() {
let r = brent(cos_minus_x, 0.0, 1.0, 1e-12, 100).unwrap();
assert_close(r, 0.739_085_133_215_160_7, 1e-12);
}
#[test]
fn brent_finds_wallis_root() {
let r = brent(wallis, 2.0, 3.0, 1e-12, 100).unwrap();
assert_close(r, 2.094_551_481_542_326_5, 1e-12);
}
#[test]
fn brent_rejects_unbracketed_interval() {
let err = brent(|x: f64| x * x, 1.0, 2.0, 1e-9, 50).unwrap_err();
assert!(matches!(err, RootError::NoBracket { .. }));
}
#[test]
fn brent_handles_root_on_endpoint() {
let r = brent(|x: f64| x - 3.0, 3.0, 4.0, 1e-12, 100).unwrap();
assert_close(r, 3.0, 1e-12);
}
#[test]
fn brent_reports_nan_evaluation() {
let err = brent(f64::ln, -1.0, 1.0, 1e-9, 50).unwrap_err();
assert!(matches!(err, RootError::NanEvaluation { .. }));
}
#[test]
fn illinois_finds_dottie_number() {
let r = illinois(cos_minus_x, 0.0, 1.0, 1e-12, 100).unwrap();
assert_close(r, 0.739_085_133_215_160_7, 1e-9);
}
#[test]
fn illinois_finds_wallis_root() {
let r = illinois(wallis, 2.0, 3.0, 1e-12, 100).unwrap();
assert_close(r, 2.094_551_481_542_326_5, 1e-9);
}
#[test]
fn illinois_rejects_unbracketed_interval() {
let err = illinois(|x: f64| x * x + 1.0, -1.0, 1.0, 1e-9, 50).unwrap_err();
assert!(matches!(err, RootError::NoBracket { .. }));
}
#[test]
fn illinois_handles_steep_function() {
let r = illinois(|x: f64| x.powi(15) + 1.0, -2.0, 0.5, 1e-9, 100).unwrap();
assert_close(r, -1.0, 1e-6);
}
}
const CGOLD: f64 = 0.381_966_011_250_105_2;
pub fn brent_minimize<F>(
mut f: F,
a: f64,
b: f64,
tol: f64,
max_iter: usize,
) -> Result<(f64, f64), RootError>
where
F: FnMut(f64) -> f64,
{
let (mut lo, mut hi) = if a < b { (a, b) } else { (b, a) };
let mut x = lo + CGOLD * (hi - lo);
let (mut w, mut v) = (x, x);
let mut fx = f(x);
let (mut fw, mut fv) = (fx, fx);
let mut d = 0.0_f64;
let mut e = 0.0_f64;
for iter in 0..max_iter {
let xm = 0.5 * (lo + hi);
let tol1 = tol * x.abs() + 1e-12;
let tol2 = 2.0 * tol1;
if (x - xm).abs() <= tol2 - 0.5 * (hi - lo) {
return Ok((x, fx));
}
let mut use_golden = true;
if e.abs() > tol1 {
let r = (x - w) * (fx - fv);
let q0 = (x - v) * (fx - fw);
let mut p = (x - v) * q0 - (x - w) * r;
let mut q = 2.0 * (q0 - r);
if q > 0.0 {
p = -p;
}
q = q.abs();
let e_temp = e;
e = d;
if p.abs() < (0.5 * q * e_temp).abs() && p > q * (lo - x) && p < q * (hi - x) {
d = p / q;
let u = x + d;
if (u - lo) < tol2 || (hi - u) < tol2 {
d = if xm - x >= 0.0 { tol1 } else { -tol1 };
}
use_golden = false;
}
}
if use_golden {
e = if x >= xm { lo - x } else { hi - x };
d = CGOLD * e;
}
let u = if d.abs() >= tol1 {
x + d
} else if d >= 0.0 {
x + tol1
} else {
x - tol1
};
let fu = f(u);
if fu <= fx {
if u >= x {
lo = x;
} else {
hi = x;
}
v = w;
fv = fw;
w = x;
fw = fx;
x = u;
fx = fu;
} else {
if u < x {
lo = u;
} else {
hi = u;
}
if fu <= fw || w == x {
v = w;
fv = fw;
w = u;
fw = fu;
} else if fu <= fv || v == x || v == w {
v = u;
fv = fu;
}
}
if iter + 1 == max_iter {
return Err(RootError::NoConvergence {
max_iter,
last_residual: hi - lo,
});
}
}
Ok((x, fx))
}
#[cfg(test)]
mod minimize_tests {
use super::*;
#[test]
fn finds_parabola_minimum() {
let (xm, fm) =
brent_minimize(|x| (x - 3.0).powi(2) + 1.0, -10.0, 10.0, 1e-10, 100).unwrap();
assert!((xm - 3.0).abs() < 1e-6, "x_min={xm}");
assert!((fm - 1.0).abs() < 1e-9, "f_min={fm}");
}
#[test]
fn finds_minimum_of_quartic() {
let f = |x: f64| x.powi(4) - 4.0 * x.powi(2) + x; let (xm, _) = brent_minimize(f, 0.5, 2.0, 1e-10, 100).unwrap();
assert!((xm - 1.34700).abs() < 1e-3, "x_min={xm}");
}
}