use multicalc::numerical_derivative::autodiff::{AutoDiffMulti, AutoDiffSingle};
use multicalc::numerical_derivative::finite_difference::FiniteDifferenceSingle;
use multicalc::root_finding::{
Bisection, Newton, NewtonSystem, RootReport, RootReportN, RootTermination,
};
use multicalc::scalar::{Numeric, ScalarFn, VectorFn, c};
use multicalc::scalar_fn;
use multicalc::scalar_fn_vec;
use multicalc::utils::error_codes::CalcError;
fn bisect<F: ScalarFn>(f: &F, a: f64, b: f64) -> Result<RootReport<f64>, CalcError> {
Bisection::default().solve(f, a, b)
}
fn newton<F: ScalarFn>(f: &F, x0: f64) -> Result<RootReport<f64>, CalcError> {
let s: Newton = Newton::default();
s.solve(f, x0)
}
fn nsystem<F: VectorFn<2, 2>>(f: &F, x0: &[f64; 2]) -> Result<RootReportN<2>, CalcError> {
let s: NewtonSystem = NewtonSystem::default();
s.solve(f, x0)
}
struct CircleHyperbola;
impl VectorFn<2, 2> for CircleHyperbola {
fn eval<S: Numeric>(&self, v: &[S; 2]) -> [S; 2] {
[c(-4.0) + v[0] * v[0] + v[1] * v[1], c(-1.0) + v[0] * v[1]]
}
}
struct TwoLinkArm {
l1: f64,
l2: f64,
px: f64,
py: f64,
}
impl VectorFn<2, 2> for TwoLinkArm {
fn eval<S: Numeric>(&self, v: &[S; 2]) -> [S; 2] {
let l1 = S::from_f64(self.l1);
let l2 = S::from_f64(self.l2);
let px = S::from_f64(self.px);
let py = S::from_f64(self.py);
[
l1 * v[0].cos() + l2 * (v[0] + v[1]).cos() - px,
l1 * v[0].sin() + l2 * (v[0] + v[1]).sin() - py,
]
}
}
#[test]
fn bisection_sqrt2() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let report = bisect(&f, 0.0_f64, 2.0).unwrap();
assert!((report.root - 2.0_f64.sqrt()).abs() < 1e-9);
assert!(matches!(
report.termination,
RootTermination::ResidualTolerance | RootTermination::BracketWidth
));
}
#[test]
fn bisection_dottie_number() {
let f = scalar_fn!(|x| x.cos() - x);
let report = bisect(&f, 0.0_f64, 1.0).unwrap();
assert!((report.root - 0.7390851332151607).abs() < 1e-9);
}
#[test]
fn bisection_wien_displacement() {
let f = scalar_fn!(|x| c(-5.0) + x + c(5.0) * (-x).exp());
let report = bisect(&f, 1.0_f64, 10.0).unwrap();
assert!((report.root - 4.965114231744276).abs() < 1e-9);
}
#[test]
fn bisection_invalid_bracket() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
assert!(matches!(
bisect(&f, 2.0_f64, 3.0),
Err(CalcError::InvalidBracket)
));
}
#[test]
fn bisection_non_finite() {
let f = scalar_fn!(|x| c(1.0) / x);
assert!(matches!(
bisect(&f, -1.0_f64, 1.0),
Err(CalcError::NonFiniteValue)
));
}
#[test]
fn bisection_budget_exhausted() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let result = Bisection::default()
.with_max_iterations(2)
.solve(&f, 0.0_f64, 2.0);
assert!(matches!(result, Err(CalcError::DidNotConverge)));
}
#[test]
fn bisection_exact_endpoint_root() {
let f = scalar_fn!(|x| x);
let report = bisect(&f, 0.0_f64, 1.0).unwrap();
assert_eq!(report.root, 0.0_f64);
assert!(matches!(
report.termination,
RootTermination::ResidualTolerance
));
}
#[test]
fn newton_sqrt2() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let report = newton(&f, 2.0_f64).unwrap();
assert!((report.root - 2.0_f64.sqrt()).abs() < 1e-12);
assert!(matches!(
report.termination,
RootTermination::ResidualTolerance | RootTermination::StepTolerance
));
}
#[test]
fn newton_cbrt2() {
let f = scalar_fn!(|x| c(-2.0) + x.powi(3));
let report = newton(&f, 1.0_f64).unwrap();
assert!((report.root - 2.0_f64.powf(1.0 / 3.0)).abs() < 1e-12);
}
#[test]
fn newton_wien_displacement() {
let f = scalar_fn!(|x| c(-5.0) + x + c(5.0) * (-x).exp());
let report = newton(&f, 5.0_f64).unwrap();
assert!((report.root - 4.965114231744276).abs() < 1e-12);
}
#[test]
fn newton_finite_difference_backend() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let solver = Newton::from_derivator(FiniteDifferenceSingle::<f64>::default());
let report = solver.solve(&f, 2.0_f64).unwrap();
assert!((report.root - 2.0_f64.sqrt()).abs() < 1e-6);
}
#[test]
fn newton_vanishing_derivative() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
assert!(matches!(
newton(&f, 0.0_f64),
Err(CalcError::SingularMatrix)
));
}
#[test]
fn newton_budget_exhausted() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let s: Newton = Newton::default().with_max_iterations(1);
assert!(matches!(
s.solve(&f, 2.0_f64),
Err(CalcError::DidNotConverge)
));
}
#[test]
fn newton_damped_rescues_far_start() {
let f = scalar_fn!(|x| x / (c(1.0) + x * x).sqrt());
let plain_result = newton(&f, 2.0_f64);
let plain_missed = match &plain_result {
Ok(r) => r.root.abs() > 0.1,
Err(_) => true,
};
assert!(
plain_missed,
"plain Newton unexpectedly converged: {plain_result:?}"
);
let damped: Newton = Newton::default().with_backtracking(true);
let report = damped.solve(&f, 2.0_f64).unwrap();
assert!(report.root.abs() < 1e-6, "{report:?}");
}
#[test]
fn newton_system_circle_hyperbola() {
let report = nsystem(&CircleHyperbola, &[1.5_f64, 0.8]).unwrap();
assert!(report.residual_norm < 1e-12);
let [x, y] = report.root;
assert!((x * x + y * y - 4.0).abs() < 1e-12);
assert!((x * y - 1.0).abs() < 1e-12);
assert!(matches!(
report.termination,
RootTermination::ResidualTolerance | RootTermination::StepTolerance
));
}
#[test]
fn newton_system_two_link_ik() {
let (l1, l2) = (1.0_f64, 1.0_f64);
let (t1_true, t2_true) = (0.5_f64, 0.8_f64);
let px = l1 * t1_true.cos() + l2 * (t1_true + t2_true).cos();
let py = l1 * t1_true.sin() + l2 * (t1_true + t2_true).sin();
let report = nsystem(&TwoLinkArm { l1, l2, px, py }, &[0.4_f64, 0.9]).unwrap();
assert!(report.residual_norm < 1e-12, "{report:?}");
let [t1, t2] = report.root;
assert!(
(t1 - t1_true).abs() < 1e-10,
"theta1: got {t1}, want {t1_true}"
);
assert!(
(t2 - t2_true).abs() < 1e-10,
"theta2: got {t2}, want {t2_true}"
);
}
#[test]
fn newton_system_singular_jacobian() {
let f = scalar_fn_vec!(|v: &[f64; 2]| [
c(-1.0) + v[0] + c(-1.0) * v[1],
c(-2.0) + c(2.0) * v[0] + c(-2.0) * v[1],
]);
assert!(matches!(
nsystem(&f, &[0.0_f64, 0.0]),
Err(CalcError::SingularMatrix)
));
}
#[test]
fn newton_system_non_finite() {
let f = scalar_fn_vec!(|v: &[f64; 2]| [c(1.0) / v[0], v[1]]);
assert!(matches!(
nsystem(&f, &[0.0_f64, 0.0]),
Err(CalcError::NonFiniteValue)
));
}
#[test]
fn newton_system_budget_exhausted() {
let s: NewtonSystem = NewtonSystem::default().with_max_iterations(1);
assert!(matches!(
s.solve(&CircleHyperbola, &[1.5_f64, 0.8]),
Err(CalcError::DidNotConverge)
));
}
#[test]
fn newton_system_damped_rescues_far_start() {
let f = scalar_fn_vec!(|v: &[f64; 2]| [
v[0] / (c(1.0) + v[0] * v[0]).sqrt(),
v[1] / (c(1.0) + v[1] * v[1]).sqrt(),
]);
let far = [3.0_f64, 3.0];
let plain_result = nsystem(&f, &far);
let plain_missed = match &plain_result {
Ok(r) => r.residual_norm > 0.1,
Err(_) => true,
};
assert!(
plain_missed,
"plain NewtonSystem unexpectedly converged: {plain_result:?}"
);
let damped: NewtonSystem = NewtonSystem::default().with_backtracking(true);
let report = damped.solve(&f, &far).unwrap();
assert!(report.residual_norm < 1e-10, "{report:?}");
}
#[test]
fn newton_sqrt2_f32() {
let f = scalar_fn!(|x| c(-2.0) + x * x);
let solver = Newton::<AutoDiffSingle<f32>>::default();
let report = solver.solve(&f, 2.0_f32).unwrap();
assert!((report.root - 2.0_f32.sqrt()).abs() < 1e-3);
}
#[test]
fn newton_system_circle_hyperbola_f32() {
let solver = NewtonSystem::<AutoDiffMulti<f32>>::default();
let report = solver.solve(&CircleHyperbola, &[1.5_f32, 0.8]).unwrap();
assert!(report.residual_norm < 1e-3);
}
struct Kepler {
e: f64,
m: f64,
}
impl ScalarFn for Kepler {
fn eval<S: Numeric>(&self, big_e: S) -> S {
big_e - S::from_f64(self.e) * big_e.sin() - S::from_f64(self.m)
}
}
#[test]
fn kepler_equation_moderate_eccentricity() {
let e = 0.8_f64;
let e_true = 1.0_f64;
let m = e_true - e * e_true.sin();
let report = newton(&Kepler { e, m }, m).unwrap();
assert!((report.root - e_true).abs() < 1e-12, "{report:?}");
assert!((report.root - e * report.root.sin() - m).abs() < 1e-12);
}
#[test]
fn kepler_equation_high_eccentricity() {
let e = 0.99_f64;
let e_true = 0.5_f64;
let m = e_true - e * e_true.sin();
let f = Kepler { e, m };
let bracketed = bisect(&f, 0.0_f64, core::f64::consts::PI).unwrap();
assert!((bracketed.root - e_true).abs() < 1e-9, "{bracketed:?}");
let damped: Newton = Newton::default().with_backtracking(true);
let stepped = damped.solve(&f, m).unwrap();
assert!((stepped.root - e_true).abs() < 1e-9, "{stepped:?}");
}
struct Colebrook {
reynolds: f64,
rel_roughness: f64,
}
impl ScalarFn for Colebrook {
fn eval<S: Numeric>(&self, f: S) -> S {
let re = S::from_f64(self.reynolds);
let eps = S::from_f64(self.rel_roughness);
let root_f = f.sqrt();
let inner = eps / S::from_f64(3.7) + S::from_f64(2.51) / (re * root_f);
let log10 = inner.ln() / S::from_f64(10.0).ln();
S::ONE / root_f + S::TWO * log10
}
}
#[test]
fn colebrook_white_friction_factor() {
let f = Colebrook {
reynolds: 1.0e5,
rel_roughness: 1.0e-4,
};
let report = newton(&f, 0.02_f64).unwrap();
assert!(report.residual.abs() < 1e-10, "{report:?}");
assert!(report.root > 0.01 && report.root < 0.05, "{report:?}");
}
struct BondYield {
cashflows: [f64; 5],
times: [i32; 5],
price: f64,
}
impl ScalarFn for BondYield {
fn eval<S: Numeric>(&self, r: S) -> S {
let mut pv = S::ZERO;
for (cash, t) in self.cashflows.iter().zip(self.times.iter()) {
pv += S::from_f64(*cash) / (S::ONE + r).powi(*t);
}
pv - S::from_f64(self.price)
}
}
#[test]
fn bond_internal_rate_of_return() {
let cashflows = [5.0, 5.0, 5.0, 5.0, 105.0];
let times = [1, 2, 3, 4, 5];
let r_true = 0.04_f64;
let price: f64 = cashflows
.iter()
.zip(times.iter())
.map(|(cash, t)| cash / (1.0_f64 + r_true).powi(*t))
.sum();
let report = newton(
&BondYield {
cashflows,
times,
price,
},
0.1_f64,
)
.unwrap();
assert!((report.root - r_true).abs() < 1e-10, "{report:?}");
}
struct Catenary {
span: f64,
length: f64,
}
impl ScalarFn for Catenary {
fn eval<S: Numeric>(&self, a: S) -> S {
let z = S::from_f64(self.span) / (S::TWO * a);
let sinh = (z.exp() - (-z).exp()) * S::HALF;
S::TWO * a * sinh - S::from_f64(self.length)
}
}
#[test]
fn catenary_constant() {
let span = 4.0_f64;
let a_true = 2.0_f64;
let z = span / (2.0 * a_true);
let length = 2.0 * a_true * z.sinh();
let report = newton(&Catenary { span, length }, 1.0_f64).unwrap();
assert!((report.root - a_true).abs() < 1e-10, "{report:?}");
}
struct DiodeLoadLine {
vs: f64,
r: f64,
is: f64,
vt: f64,
}
impl ScalarFn for DiodeLoadLine {
fn eval<S: Numeric>(&self, v: S) -> S {
let vs = S::from_f64(self.vs);
let r = S::from_f64(self.r);
let is = S::from_f64(self.is);
let vt = S::from_f64(self.vt);
(vs - v) / r - is * ((v / vt).exp() - S::ONE)
}
}
#[test]
fn diode_load_line_voltage() {
let vs = 5.0_f64;
let r = 1000.0_f64;
let vt = 0.025852_f64;
let v_true = 0.6_f64;
let is = ((vs - v_true) / r) / ((v_true / vt).exp() - 1.0);
let diode = DiodeLoadLine { vs, r, is, vt };
let report = bisect(&diode, 0.0_f64, 1.0).unwrap();
assert!((report.root - v_true).abs() < 1e-9, "{report:?}");
}
#[test]
fn wien_displacement_constant_from_blackbody_peak() {
let f = scalar_fn!(|x| c(-5.0) + x + c(5.0) * (-x).exp());
let report = newton(&f, 5.0_f64).unwrap();
let x = report.root;
let h = 6.62607015e-34_f64;
let c_light = 299_792_458.0_f64;
let k_b = 1.380649e-23_f64;
let b = h * c_light / (x * k_b);
assert!((b - 2.897771955e-3).abs() < 1e-9, "b = {b}");
}
#[test]
fn chemical_equilibrium_three_species() {
let f = scalar_fn_vec!(|v: &[f64; 3]| [
c(-1.0) + v[0] + v[1] + v[2],
v[1] - c(1.25) * v[0] * v[0],
v[2] - c(5.0) * v[0] * v[1],
]);
let solver: NewtonSystem = NewtonSystem::default();
let report = solver.solve(&f, &[0.5_f64, 0.25, 0.25]).unwrap();
assert!(report.residual_norm < 1e-12, "{report:?}");
let [a, b, conc_c] = report.root;
assert!((a - 0.4).abs() < 1e-10, "{report:?}");
assert!((b - 0.2).abs() < 1e-10, "{report:?}");
assert!((conc_c - 0.4).abs() < 1e-10, "{report:?}");
}
#[test]
fn two_link_arm_far_start_damped() {
let (l1, l2) = (2.0_f64, 1.0_f64);
let (t1_true, t2_true) = (0.6_f64, 0.9_f64);
let px = l1 * t1_true.cos() + l2 * (t1_true + t2_true).cos();
let py = l1 * t1_true.sin() + l2 * (t1_true + t2_true).sin();
let arm = TwoLinkArm { l1, l2, px, py };
let solver: NewtonSystem = NewtonSystem::default().with_backtracking(true);
let report = solver.solve(&arm, &[0.1_f64, 0.5]).unwrap();
assert!(report.residual_norm < 1e-10, "{report:?}");
let [t1, t2] = report.root;
let tip_x = l1 * t1.cos() + l2 * (t1 + t2).cos();
let tip_y = l1 * t1.sin() + l2 * (t1 + t2).sin();
assert!(
(tip_x - px).abs() < 1e-9 && (tip_y - py).abs() < 1e-9,
"{report:?}"
);
}