spectrafit_newton_cg/
step.rs1use faer::Mat;
19use spectrafit_trust_region::{StepResult, Subproblem, SubproblemStep};
20
21pub struct SteihaugStep;
23
24#[inline]
25fn dot(a: &Mat<f64>, b: &Mat<f64>) -> f64 {
26 (a.as_ref().transpose() * b.as_ref())[(0, 0)]
27}
28
29#[inline]
30fn norm(v: &Mat<f64>) -> f64 {
31 v.as_ref().squared_norm_l2().sqrt()
32}
33
34fn boundary_tau(z: &Mat<f64>, d: &Mat<f64>, radius: f64) -> f64 {
36 let a = dot(d, d);
37 if a <= 0.0 {
38 return 0.0;
39 }
40 let b = 2.0 * dot(z, d);
41 let c = dot(z, z) - radius * radius;
42 let disc = (b * b - 4.0 * a * c).max(0.0);
43 ((-b + disc.sqrt()) / (2.0 * a)).max(0.0)
44}
45
46impl SubproblemStep for SteihaugStep {
47 fn solve(&self, sub: &Subproblem<'_>, radius: f64) -> StepResult {
48 let p = sub.n_params();
49 let g = sub.gradient(); let finish = |step: Mat<f64>, hit_boundary: bool| {
52 let predicted_reduction = sub.predicted_reduction(step.as_ref());
53 StepResult {
54 step,
55 predicted_reduction,
56 hit_boundary,
57 }
58 };
59
60 let mut z = Mat::<f64>::zeros(p, 1);
62 let mut r = Mat::from_fn(p, 1, |i, _| g[(i, 0)]);
63 let r0_norm = norm(&r);
64 if r0_norm == 0.0 {
65 return finish(z, false); }
67 let tol = r0_norm * 0.5_f64.min(r0_norm.sqrt());
69 let mut d = Mat::from_fn(p, 1, |i, _| -r[(i, 0)]);
70 let mut rr = dot(&r, &r);
71
72 let max_iter = 2 * p + 10;
73 for _ in 0..max_iter {
74 let hd = sub.hvec(d.as_ref());
75 let dhd = dot(&d, &hd);
76 if dhd <= 0.0 {
77 let tau = boundary_tau(&z, &d, radius);
79 let step = Mat::from_fn(p, 1, |i, _| z[(i, 0)] + tau * d[(i, 0)]);
80 return finish(step, true);
81 }
82 let alpha = rr / dhd;
83 let z_next = Mat::from_fn(p, 1, |i, _| z[(i, 0)] + alpha * d[(i, 0)]);
84 if norm(&z_next) >= radius {
85 let tau = boundary_tau(&z, &d, radius);
87 let step = Mat::from_fn(p, 1, |i, _| z[(i, 0)] + tau * d[(i, 0)]);
88 return finish(step, true);
89 }
90 z = z_next;
91 let r_next = Mat::from_fn(p, 1, |i, _| r[(i, 0)] + alpha * hd[(i, 0)]);
92 if norm(&r_next) <= tol {
93 return finish(z, false); }
95 let rr_next = dot(&r_next, &r_next);
96 let beta = rr_next / rr;
97 d = Mat::from_fn(p, 1, |i, _| -r_next[(i, 0)] + beta * d[(i, 0)]);
98 r = r_next;
99 rr = rr_next;
100 }
101 finish(z, false) }
103}