Skip to main content

gam_model_kernels/
monotone_root.rs

1/// Shared hybrid bracketing + Newton solver for strictly monotone calibration
2/// equations F(a) = 0.
3///
4/// `eval(a)` must return `(F(a), F'(a), F''(a))`.  The second derivative feeds
5/// the safeguarded Halley step inside the shared solver.
6///
7/// The algorithm core lives in `opt::find_root_monotone` (warm-start Newton
8/// probes → analytic-or-geometric bracket → hybrid Halley/Newton/bisection
9/// refinement with best-point tracking); this module is the thin domain
10/// wrapper that keeps the calibration-facing signature, threads the `label`
11/// context through, and maps `opt::RootError` back onto
12/// [`MonotoneRootError`]'s byte-identical `Display` shapes.
13///
14/// Returns `(root, |F'(root)|, F(root))`.  The absolute derivative is always
15/// positive and can be used directly as the density-normalising calibration
16/// derivative.  Callers must validate the residual against the scale of their
17/// calibration equation.
18///
19/// The monotone direction (increasing vs decreasing) is inferred from the
20/// sign of F'(a) at the initial point, so the same code handles both the
21/// Bernoulli case (F increasing) and the survival case (F decreasing).
22#[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    // `RootConfig::new` fills exactly this solver's historical constants:
80    // bracket step fraction 0.25, Newton/Halley derivative floor 1e-30, and
81    // the ×8 warm-start trust cap.
82    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
98/// Fold the shared solver's typed error back onto this crate's
99/// [`MonotoneRootError`], preserving the exact pre-refactor `Display`
100/// strings (each arm routes through the same factory the old in-crate
101/// solver used at the corresponding failure site).
102fn 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        // F(a) = 2a − 7, root at a = 3.5, F' = 2 everywhere.
199        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        // F(0) = 0 exactly, so the solver should return immediately with refine_iters=0.
216        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        // F'(a_init) = 0 is degenerate; the solver must return Err rather than
232        // infinite-loop or divide by zero.
233        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        // Supply a bracket [0, 10] for F(a) = a − 3; the solver must use it and
251        // converge to root = 3.
252        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        // max_bracket_iters=0 and init is not at root → bracketing cannot succeed.
269        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, // no bracket iterations allowed
275            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        // The Display strings are contract: they must be byte-identical to the
287        // pre-extraction in-crate solver's output at each failure site.
288        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}