1#[derive(Clone, Copy, Debug)]
23pub struct MonotoneRootSolution {
24 pub root: f64,
25 pub abs_deriv: f64,
26 pub residual: f64,
27 pub refine_iters: usize,
28}
29
30pub use gam_problem::MonotoneRootError;
31
32pub fn solve_monotone_root(
33 eval: impl Fn(f64) -> Result<(f64, f64, f64), String>,
34 a_init: f64,
35 label: &str,
36 convergence_tol: f64,
37 max_bracket_iters: usize,
38 max_refine_iters: usize,
39) -> Result<(f64, f64, f64), MonotoneRootError> {
40 let solution = solve_monotone_root_detailed(
41 eval,
42 a_init,
43 label,
44 convergence_tol,
45 max_bracket_iters,
46 max_refine_iters,
47 )?;
48 Ok((solution.root, solution.abs_deriv, solution.residual))
49}
50
51pub fn solve_monotone_root_detailed(
52 eval: impl Fn(f64) -> Result<(f64, f64, f64), String>,
53 a_init: f64,
54 label: &str,
55 convergence_tol: f64,
56 max_bracket_iters: usize,
57 max_refine_iters: usize,
58) -> Result<MonotoneRootSolution, MonotoneRootError> {
59 solve_monotone_root_detailed_with_bracket(
60 eval,
61 a_init,
62 label,
63 convergence_tol,
64 max_bracket_iters,
65 max_refine_iters,
66 None,
67 )
68}
69
70pub fn solve_monotone_root_detailed_with_bracket(
71 eval: impl Fn(f64) -> Result<(f64, f64, f64), String>,
72 a_init: f64,
73 label: &str,
74 convergence_tol: f64,
75 max_bracket_iters: usize,
76 max_refine_iters: usize,
77 analytic_bracket: Option<(f64, f64)>,
78) -> Result<MonotoneRootSolution, MonotoneRootError> {
79 let config = opt::RootConfig::new(convergence_tol, max_bracket_iters, max_refine_iters);
83 let oracle = |a: f64| {
84 eval(a)
85 .map(|(value, d1, d2)| opt::RootSample { value, d1, d2 })
86 .map_err(|source| (a, source))
87 };
88 let solution = opt::find_root_monotone(oracle, a_init, &config, analytic_bracket)
89 .map_err(|err| map_root_err(label, err))?;
90 Ok(MonotoneRootSolution {
91 root: solution.root,
92 abs_deriv: solution.abs_deriv,
93 residual: solution.value,
94 refine_iters: solution.iters,
95 })
96}
97
98fn map_root_err(label: &str, err: opt::RootError<(f64, String)>) -> MonotoneRootError {
103 match err {
104 opt::RootError::Eval((a, source)) => MonotoneRootError::EvalFailed {
105 label: label.to_string(),
106 a,
107 source,
108 },
109 opt::RootError::ExactRootDegenerate { at } => {
110 MonotoneRootError::exact_root_degenerate(label, at)
111 }
112 opt::RootError::DegenerateDerivative { at, derivative } => {
113 MonotoneRootError::DegenerateDerivative {
114 label: label.to_string(),
115 a: at,
116 fp: derivative,
117 }
118 }
119 opt::RootError::BracketInvalid { lo, hi } => {
120 MonotoneRootError::analytic_bracket_invalid(label, lo, hi)
121 }
122 opt::RootError::BracketNoStraddle { f_lo, f_hi } => {
123 MonotoneRootError::analytic_bracket_no_straddle(label, f_lo, f_hi)
124 }
125 opt::RootError::BracketingExhausted { direction, seed } => {
126 MonotoneRootError::search_exhausted(label, direction, seed)
127 }
128 opt::RootError::ConvergedRootDegenerate { at } => {
129 MonotoneRootError::converged_root_degenerate(label, at)
130 }
131 }
132}
133
134#[cfg(test)]
135mod tests {
136 use super::{
137 MonotoneRootError, solve_monotone_root, solve_monotone_root_detailed,
138 solve_monotone_root_detailed_with_bracket,
139 };
140 use std::cell::RefCell;
141
142 #[test]
143 fn solve_monotone_root_converges_for_increasing_function() {
144 let (root, abs_deriv, residual) = solve_monotone_root(
145 |a| {
146 let ea = a.exp();
147 Ok((ea - 2.0, ea, ea))
148 },
149 0.0,
150 "increasing",
151 1e-12,
152 32,
153 32,
154 )
155 .expect("root");
156
157 assert!((root - std::f64::consts::LN_2).abs() < 1e-10);
158 assert!((abs_deriv - 2.0).abs() < 1e-10);
159 assert!(residual.abs() < 1e-12);
160 }
161
162 #[test]
163 fn solve_monotone_root_accepts_halley_probe_for_decreasing_function() {
164 let eval_points = RefCell::new(Vec::new());
165 let (root, abs_deriv, residual) = solve_monotone_root(
166 |a| {
167 eval_points.borrow_mut().push(a);
168 let ea = (-a).exp();
169 Ok((ea - 0.5, -ea, ea))
170 },
171 0.0,
172 "decreasing",
173 1e-12,
174 32,
175 32,
176 )
177 .expect("root");
178
179 let f_mid = (-0.5f64).exp() - 0.5;
180 let f_a_mid = -(-0.5f64).exp();
181 let f_aa_mid = (-0.5f64).exp();
182 let expected_probe =
183 0.5 - (2.0 * f_mid * f_a_mid) / (2.0 * f_a_mid * f_a_mid - f_mid * f_aa_mid);
184 assert!((root - std::f64::consts::LN_2).abs() < 1e-10);
185 assert!((abs_deriv - 0.5).abs() < 1e-10);
186 assert!(residual.abs() < 1e-12);
187 assert!(
188 eval_points
189 .borrow()
190 .iter()
191 .copied()
192 .any(|a| (a - expected_probe).abs() < 1e-12)
193 );
194 }
195
196 #[test]
197 fn solve_linear_function_reaches_exact_root() {
198 let (root, abs_deriv, residual) = solve_monotone_root(
200 |a| Ok((2.0 * a - 7.0, 2.0, 0.0)),
201 0.0,
202 "linear",
203 1e-12,
204 32,
205 64,
206 )
207 .expect("root");
208 assert!((root - 3.5).abs() < 1e-10, "root={root}");
209 assert!((abs_deriv - 2.0).abs() < 1e-10, "abs_deriv={abs_deriv}");
210 assert!(residual.abs() < 1e-12, "residual={residual}");
211 }
212
213 #[test]
214 fn exact_root_at_init_returns_zero_iters() {
215 let result = solve_monotone_root_detailed(
217 |a| Ok((a, 1.0, 0.0)),
218 0.0,
219 "exact_at_init",
220 1e-12,
221 32,
222 32,
223 )
224 .expect("solution");
225 assert!(result.root.abs() < 1e-12, "root={}", result.root);
226 assert_eq!(result.refine_iters, 0);
227 }
228
229 #[test]
230 fn degenerate_derivative_returns_error() {
231 let err = solve_monotone_root(
234 |a| Ok((a - 5.0, 0.0, 0.0)),
235 0.0,
236 "degenerate_fp",
237 1e-12,
238 32,
239 32,
240 )
241 .unwrap_err();
242 match err {
243 MonotoneRootError::DegenerateDerivative { .. } => {}
244 other => panic!("expected DegenerateDerivative, got {other:?}"),
245 }
246 }
247
248 #[test]
249 fn analytic_bracket_is_honored() {
250 let sol = solve_monotone_root_detailed_with_bracket(
253 |a| Ok((a - 3.0, 1.0, 0.0)),
254 5.0,
255 "analytic_bracket",
256 1e-12,
257 32,
258 64,
259 Some((0.0, 10.0)),
260 )
261 .expect("solution");
262 assert!((sol.root - 3.0).abs() < 1e-10, "root={}", sol.root);
263 assert!(sol.residual.abs() < 1e-12, "residual={}", sol.residual);
264 }
265
266 #[test]
267 fn search_exhausted_with_zero_bracket_iters() {
268 let err = solve_monotone_root(
270 |a| Ok((a - 100.0, 1.0, 0.0)),
271 0.0,
272 "no_bracket",
273 1e-12,
274 0, 32,
276 )
277 .unwrap_err();
278 match err {
279 MonotoneRootError::BracketingExhausted { .. } => {}
280 other => panic!("expected BracketingExhausted, got {other:?}"),
281 }
282 }
283
284 #[test]
285 fn error_display_shapes_are_preserved_through_the_shared_solver() {
286 let degenerate =
289 solve_monotone_root(|a| Ok((a - 5.0, 0.0, 0.0)), 0.0, "caldbg", 1e-12, 32, 32)
290 .unwrap_err();
291 assert_eq!(
292 degenerate.to_string(),
293 "caldbg: initial derivative is zero or non-finite at a=0.000000"
294 );
295
296 let exhausted =
297 solve_monotone_root(|a| Ok((a - 100.0, 1.0, 0.0)), 0.0, "caldbg", 1e-12, 0, 32)
298 .unwrap_err();
299 assert_eq!(
300 exhausted.to_string(),
301 "caldbg: failed to bracket root (searched +1 from a=0.000000)"
302 );
303
304 let eval_failed = solve_monotone_root(
305 |_| Err("inner calibration blew up".to_string()),
306 0.0,
307 "caldbg",
308 1e-12,
309 32,
310 32,
311 )
312 .unwrap_err();
313 assert_eq!(eval_failed.to_string(), "inner calibration blew up");
314 }
315}