Skip to main content

diffsol/ode_solver/
mod.rs

1pub mod adjoint;
2pub mod bdf;
3pub mod bdf_state;
4pub mod builder;
5pub mod checkpointing;
6pub mod config;
7pub mod explicit_rk;
8pub mod jacobian_update;
9pub mod method;
10pub mod no_checkpointing_solver;
11pub mod problem;
12pub mod runge_kutta;
13pub mod sde;
14pub mod sdirk;
15pub mod sdirk_state;
16pub mod sensitivities;
17pub mod solution;
18pub mod state;
19pub mod tableau;
20
21use serde::Serialize;
22use std::fmt::Display;
23
24use crate::ode_solver::jacobian_update::SolverState;
25
26/// Solver statistics shared by all ODE solver methods.
27#[derive(Clone, Debug, Serialize, Default)]
28pub struct OdeSolverStatistics {
29    /// Total Jacobian/LU setups (CVODE `nsetups`); sum of the per-cause counters below.
30    pub number_of_linear_solver_setups: usize,
31    /// Number of time steps taken by the solver.
32    pub number_of_steps: usize,
33    /// Number of local error test failures (steps rejected for excessive local error).
34    pub number_of_error_test_failures: usize,
35    /// Total number of nonlinear (Newton) solver iterations across all steps.
36    pub number_of_nonlinear_solver_iterations: usize,
37    /// Number of nonlinear (Newton) solver convergence failures.
38    pub number_of_nonlinear_solver_fails: usize,
39    /// Jacobian/LU setups triggered by checkpoint or reinitialisation.
40    pub number_of_linear_solver_setups_from_checkpoint: usize,
41    /// Jacobian/LU setups triggered by a first nonlinear convergence failure.
42    pub number_of_linear_solver_setups_from_first_convergence_fail: usize,
43    /// Jacobian/LU setups triggered by a second nonlinear convergence failure.
44    pub number_of_linear_solver_setups_from_second_convergence_fail: usize,
45    /// Jacobian/LU setups triggered by a local error test failure.
46    pub number_of_linear_solver_setups_from_error_test_fail: usize,
47    /// Jacobian/LU setups triggered by the normal step-success heuristic.
48    pub number_of_linear_solver_setups_from_step_success: usize,
49}
50
51impl OdeSolverStatistics {
52    /// Record a Jacobian/LU setup, incrementing the total and the per-cause counter.
53    pub(crate) fn record_linear_solver_setup(&mut self, cause: SolverState) {
54        self.number_of_linear_solver_setups += 1;
55        match cause {
56            SolverState::Checkpoint => self.number_of_linear_solver_setups_from_checkpoint += 1,
57            SolverState::FirstConvergenceFail => {
58                self.number_of_linear_solver_setups_from_first_convergence_fail += 1
59            }
60            SolverState::SecondConvergenceFail => {
61                self.number_of_linear_solver_setups_from_second_convergence_fail += 1
62            }
63            SolverState::ErrorTestFail => {
64                self.number_of_linear_solver_setups_from_error_test_fail += 1
65            }
66            SolverState::StepSuccess => self.number_of_linear_solver_setups_from_step_success += 1,
67        }
68    }
69}
70
71impl Display for OdeSolverStatistics {
72    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
73        write!(f, "{}", serde_json::to_string_pretty(self).unwrap())
74    }
75}
76
77#[cfg(test)]
78mod tests {
79    use std::rc::Rc;
80
81    use self::problem::OdeSolverSolution;
82
83    use super::*;
84    use crate::error::{DiffsolError, OdeSolverError};
85    use crate::matrix::Matrix;
86    use crate::ode_solver::sensitivities::SensitivitiesOdeSolverMethod;
87    use crate::ode_solver::solution::Solution;
88    use crate::op::unit::UnitCallable;
89    use crate::op::ParameterisedOp;
90    use crate::Scalar;
91    use crate::{
92        op::OpStatistics, AdjointEquations, AdjointOdeSolverMethod, Context, DenseMatrix,
93        MatrixCommon, MatrixRef, NonLinearOp, NonLinearOpJacobian, OdeEquations,
94        OdeEquationsImplicit, OdeEquationsImplicitAdjoint, OdeEquationsImplicitSens,
95        OdeEquationsRef, OdeSolverConfig, OdeSolverMethod, OdeSolverProblem, OdeSolverState,
96        OdeSolverStopReason, Scale, VectorRef, VectorView, VectorViewMut,
97    };
98    use crate::{
99        ConstantOp, ConstantOpSens, DefaultDenseMatrix, DefaultSolver, LinearSolver,
100        NonLinearOpSens, Op, Vector,
101    };
102    use num_traits::{FromPrimitive, One, Signed, ToPrimitive, Zero};
103
104    pub fn test_ode_solver<'a, M, Eqn, Method>(
105        method: &mut Method,
106        solution: OdeSolverSolution<M::V>,
107        override_tol: Option<M::T>,
108        use_tstop: bool,
109        solve_for_sensitivities: bool,
110    ) -> Eqn::V
111    where
112        M: Matrix,
113        Eqn: OdeEquations<M = M, T = M::T, V = M::V> + 'a,
114        Method: OdeSolverMethod<'a, Eqn>,
115    {
116        let have_root = method.problem().eqn.root().is_some();
117        for (i, point) in solution.solution_points.iter().enumerate() {
118            let (soln, sens_soln) = if use_tstop {
119                match method.set_stop_time(point.t) {
120                    Ok(_) => loop {
121                        match method.step() {
122                            Ok(OdeSolverStopReason::RootFound(_, _)) => {
123                                assert!(have_root);
124                                return method.state().y.clone();
125                            }
126                            Ok(OdeSolverStopReason::TstopReached) => {
127                                break (method.state().y.clone(), method.state().s.to_vec());
128                            }
129                            _ => (),
130                        }
131                    },
132                    Err(_) => (method.state().y.clone(), method.state().s.to_vec()),
133                }
134            } else {
135                while method.state().t.abs() < point.t.abs() {
136                    if let OdeSolverStopReason::RootFound(t, _) = method.step().unwrap() {
137                        assert!(have_root);
138                        return method.interpolate(t).unwrap();
139                    }
140                }
141                let soln = method.interpolate(point.t).unwrap();
142                let sens_soln = method.interpolate_sens(point.t).unwrap();
143                (soln, sens_soln)
144            };
145            let soln = if let Some(out) = method.problem().eqn.out() {
146                out.call(&soln, point.t)
147            } else {
148                soln
149            };
150            assert_eq!(
151                soln.len(),
152                point.state.len(),
153                "soln.len() != point.state.len()"
154            );
155            if let Some(override_tol) = override_tol {
156                soln.assert_eq_st(&point.state, override_tol);
157            } else {
158                let (rtol, atol) = if method.problem().eqn.out().is_some() {
159                    // problem rtol and atol is on the state, so just use solution tolerance here
160                    (solution.rtol, &solution.atol)
161                } else {
162                    (method.problem().rtol, &method.problem().atol)
163                };
164                let error = soln.clone() - &point.state;
165                let error_norm = error.squared_norm(&point.state, atol, rtol).sqrt();
166                assert!(
167                    error_norm < M::T::from_f64(20.0).unwrap(),
168                    "error_norm: {} at t = {}. soln: {:?}, expected: {:?}",
169                    error_norm,
170                    point.t,
171                    soln,
172                    point.state
173                );
174                if solve_for_sensitivities {
175                    if let Some(sens_soln_points) = solution.sens_solution_points.as_ref() {
176                        for (j, sens_points) in sens_soln_points.iter().enumerate() {
177                            let sens_point = &sens_points[i];
178                            let sens_soln = &sens_soln[j];
179                            let error = sens_soln.clone() - &sens_point.state;
180                            let error_norm =
181                                error.squared_norm(&sens_point.state, atol, rtol).sqrt();
182                            assert!(
183                                error_norm < M::T::from_f64(29.0).unwrap(),
184                                "error_norm: {error_norm} at t = {}, sens index: {j}. soln: {sens_soln:?}, expected: {:?}",
185                                point.t,
186                                sens_point.state
187                            );
188                        }
189                    }
190                }
191            }
192        }
193        method.state().y.clone()
194    }
195
196    pub fn setup_test_adjoint<'a, LS, Eqn>(
197        problem: &'a mut OdeSolverProblem<Eqn>,
198        soln: OdeSolverSolution<Eqn::V>,
199    ) -> <Eqn::V as DefaultDenseMatrix>::M
200    where
201        Eqn: OdeEquationsImplicitAdjoint + 'a,
202        LS: LinearSolver<Eqn::M>,
203        Eqn::M: DefaultSolver,
204        Eqn::V: DefaultDenseMatrix,
205        for<'b> &'b Eqn::V: VectorRef<Eqn::V>,
206        for<'b> &'b Eqn::M: MatrixRef<Eqn::M>,
207    {
208        let nparams = problem.eqn.nparams();
209        let nout = problem.eqn.nout();
210        let ctx = problem.eqn.context();
211        let mut dgdp = <Eqn::V as DefaultDenseMatrix>::M::zeros(nparams, nout, ctx.clone());
212        let final_time = soln.solution_points.last().unwrap().t;
213        let mut p_0 = Eqn::V::zeros(nparams, ctx.clone());
214        problem.eqn.get_params(&mut p_0);
215        let nbatch = p_0.context().nbatch();
216        let h_base = Eqn::T::from_f64(1e-6).unwrap();
217        let mut h = Eqn::V::from_element(nparams, h_base, ctx.clone());
218        h.axpy(h_base, &p_0, Eqn::T::one());
219        let p_base = p_0.clone();
220        for i in 0..nparams {
221            for b in 0..nbatch {
222                let base = p_base.get_batch(b).get_index(i);
223                let hb = h.get_batch(b).get_index(i);
224                p_0.get_batch_mut(b).set_index(i, base + hb);
225            }
226            problem.eqn.set_params(&p_0);
227            let g_pos = {
228                let mut s = problem.bdf::<LS>().unwrap();
229                s.solve(final_time).unwrap();
230                s.state().g.clone()
231            };
232
233            for b in 0..nbatch {
234                let base = p_base.get_batch(b).get_index(i);
235                let hb = h.get_batch(b).get_index(i);
236                p_0.get_batch_mut(b).set_index(i, base - hb);
237            }
238            problem.eqn.set_params(&p_0);
239            let g_neg = {
240                let mut s = problem.bdf::<LS>().unwrap();
241                s.solve(final_time).unwrap();
242                s.state().g.clone()
243            };
244            for b in 0..nbatch {
245                let base = p_base.get_batch(b).get_index(i);
246                p_0.get_batch_mut(b).set_index(i, base);
247            }
248
249            let delta_full = g_pos - g_neg;
250            for b in 0..nbatch {
251                let hb = h.get_batch(b).get_index(i);
252                let denom = Eqn::T::from_f64(2.0).unwrap() * hb;
253                for j in 0..nout {
254                    let delta_val = delta_full.get_batch(b).get_index(j) / denom;
255                    dgdp.set_index(i, b * nout + j, delta_val);
256                }
257            }
258        }
259        problem.eqn.set_params(&p_base);
260        dgdp
261    }
262
263    /// sum_i^n (soln_i - data_i)^2
264    /// sum_i^n (soln_i - data_i)^4
265    pub(crate) fn sum_squares<DM>(soln: &DM, data: &DM) -> DM::V
266    where
267        DM: DenseMatrix,
268    {
269        let nbatch = soln.context().nbatch();
270        let mut ret = DM::V::zeros(2, soln.context().clone());
271        for j in 0..soln.ncols() {
272            let soln_j = soln.column(j);
273            let data_j = data.column(j);
274            let delta = soln_j - data_j;
275            for b in 0..nbatch {
276                let delta_b = delta.get_batch(b).into_owned();
277                let norm2 = delta_b.norm(2);
278                let norm4 = delta_b.norm(4);
279                let cur0 = ret.get_batch(b).get_index(0);
280                let cur1 = ret.get_batch(b).get_index(1);
281                ret.get_batch_mut(b).set_index(0, cur0 + norm2 * norm2);
282                let norm4_sq = norm4 * norm4;
283                ret.get_batch_mut(b)
284                    .set_index(1, cur1 + norm4_sq * norm4_sq);
285            }
286        }
287        ret
288    }
289
290    /// sum_i^n 2 * (soln_i - data_i)
291    /// sum_i^n 4 * (soln_i - data_i)^3
292    pub(crate) fn dsum_squaresdp<DM>(soln: &DM, data: &DM) -> Vec<DM>
293    where
294        DM: DenseMatrix,
295    {
296        let delta = soln.clone() - data;
297        let mut delta3 = delta.clone();
298        for j in 0..delta3.ncols() {
299            let delta_col = delta.column(j).into_owned();
300
301            let mut delta3_col = delta_col.clone();
302            delta3_col.component_mul_assign(&delta_col);
303            delta3_col.component_mul_assign(&delta_col);
304
305            delta3.column_mut(j).copy_from(&delta3_col);
306        }
307        let ret = vec![
308            delta * Scale(DM::T::from_f64(2.).unwrap()),
309            delta3 * Scale(DM::T::from_f64(4.).unwrap()),
310        ];
311        ret
312    }
313
314    pub fn setup_test_adjoint_sum_squares<'a, LS, Eqn>(
315        problem: &'a mut OdeSolverProblem<Eqn>,
316        times: &[Eqn::T],
317    ) -> (
318        <Eqn::V as DefaultDenseMatrix>::M,
319        <Eqn::V as DefaultDenseMatrix>::M,
320    )
321    where
322        Eqn: OdeEquationsImplicitAdjoint + 'a,
323        LS: LinearSolver<Eqn::M>,
324        Eqn::M: DefaultSolver,
325        Eqn::V: DefaultDenseMatrix,
326        for<'b> &'b Eqn::V: VectorRef<Eqn::V>,
327        for<'b> &'b Eqn::M: MatrixRef<Eqn::M>,
328    {
329        let nparams = problem.eqn.nparams();
330        let nout = 2;
331        let ctx = problem.eqn.context();
332        let mut dgdp = <Eqn::V as DefaultDenseMatrix>::M::zeros(nparams, nout, ctx.clone());
333
334        let mut p_0 = ctx.vector_zeros(nparams);
335        problem.eqn.get_params(&mut p_0);
336        let nbatch = p_0.context().nbatch();
337        let h_base = Eqn::T::from_f64(1e-6).unwrap();
338        let mut h = Eqn::V::from_element(nparams, h_base, ctx.clone());
339        h.axpy(h_base, &p_0, Eqn::T::one());
340        let mut p_data = p_0.clone();
341        p_data.axpy(Eqn::T::from_f64(0.1).unwrap(), &p_0, Eqn::T::one());
342        let p_base = p_0.clone();
343
344        problem.eqn.set_params(&p_data);
345        let data = {
346            let mut s = problem.bdf::<LS>().unwrap();
347            s.solve_dense(times).unwrap().0
348        };
349
350        for i in 0..nparams {
351            for b in 0..nbatch {
352                let base = p_base.get_batch(b).get_index(i);
353                let hb = h.get_batch(b).get_index(i);
354                p_0.get_batch_mut(b).set_index(i, base + hb);
355            }
356            problem.eqn.set_params(&p_0);
357            let g_pos = {
358                let mut s = problem.bdf::<LS>().unwrap();
359                let v = s.solve_dense(times).unwrap().0;
360                sum_squares(&v, &data)
361            };
362
363            for b in 0..nbatch {
364                let base = p_base.get_batch(b).get_index(i);
365                let hb = h.get_batch(b).get_index(i);
366                p_0.get_batch_mut(b).set_index(i, base - hb);
367            }
368            problem.eqn.set_params(&p_0);
369            let g_neg = {
370                let mut s = problem.bdf::<LS>().unwrap();
371                let v = s.solve_dense(times).unwrap().0;
372                sum_squares(&v, &data)
373            };
374
375            for b in 0..nbatch {
376                let base = p_base.get_batch(b).get_index(i);
377                p_0.get_batch_mut(b).set_index(i, base);
378            }
379
380            let delta_full = g_pos - g_neg;
381            for b in 0..nbatch {
382                let hb = h.get_batch(b).get_index(i);
383                let denom = Eqn::T::from_f64(2.0).unwrap() * hb;
384                for j in 0..nout {
385                    let delta_val = delta_full.get_batch(b).get_index(j) / denom;
386                    dgdp.set_index(i, b * nout + j, delta_val);
387                }
388            }
389        }
390        problem.eqn.set_params(&p_base);
391        (dgdp, data)
392    }
393
394    pub fn single_reset_root_discrete_times<T: Scalar>(t_stop: T) -> Vec<T> {
395        let t_root = t_stop / T::from_f64(2.0).unwrap();
396        [0.25, 0.75, 1.25, 1.75]
397            .into_iter()
398            .map(|factor| t_root * T::from_f64(factor).unwrap())
399            .collect()
400    }
401
402    fn solve_dense_with_single_reset_root<'a, Eqn, Method, BuildForward>(
403        build_forward: BuildForward,
404        times: &[Eqn::T],
405    ) -> <Eqn::V as DefaultDenseMatrix>::M
406    where
407        Eqn: OdeEquationsImplicitAdjoint + 'a,
408        Eqn::M: DefaultSolver,
409        Eqn::V: DefaultDenseMatrix,
410        Method: OdeSolverMethod<'a, Eqn>,
411        BuildForward: Fn(Option<Method::State>) -> Result<Method, DiffsolError>,
412    {
413        let mut soln = Solution::<Eqn::V>::new_dense(times.to_vec()).unwrap();
414        let first_forward_solver = build_forward(None).unwrap().solve_soln(&mut soln).unwrap();
415        match soln.stop_reason {
416            Some(OdeSolverStopReason::RootFound(_, 0)) => {}
417            Some(OdeSolverStopReason::RootFound(_, idx)) => {
418                panic!("expected first solve_soln() segment to stop on root 0, got root {idx}")
419            }
420            Some(OdeSolverStopReason::TstopReached) => {
421                panic!("expected first solve_soln() segment to stop on the interior root")
422            }
423            Some(OdeSolverStopReason::InternalTimestep) | None => {
424                panic!("first solve_soln() segment did not finish with a terminal stop reason")
425            }
426        }
427
428        let mut state_after_reset = first_forward_solver.state_clone();
429        {
430            let problem = first_forward_solver.problem();
431            state_after_reset
432                .as_mut()
433                .apply_reset_with_mass::<<Eqn::M as DefaultSolver>::LS, _>(problem)
434                .unwrap();
435        }
436
437        build_forward(Some(state_after_reset))
438            .unwrap()
439            .solve_soln(&mut soln)
440            .unwrap();
441        assert!(
442            soln.is_complete(),
443            "expected stitched solve_soln() output to cover all requested observation times",
444        );
445        soln.ys
446    }
447
448    fn state_after_manual_reset<'a, Eqn, Method>(solver: &Method) -> Method::State
449    where
450        Eqn: OdeEquationsImplicitAdjoint + 'a,
451        Eqn::M: DefaultSolver,
452        Method: OdeSolverMethod<'a, Eqn>,
453    {
454        let mut state_after_reset = solver.state_clone();
455        {
456            let problem = solver.problem();
457            state_after_reset
458                .as_mut()
459                .apply_reset_with_mass::<<Eqn::M as DefaultSolver>::LS, _>(problem)
460                .unwrap();
461        }
462        state_after_reset
463    }
464
465    pub fn setup_test_adjoint_sum_squares_with_single_reset_root<'a, LS, Eqn>(
466        problem: &'a mut OdeSolverProblem<Eqn>,
467        times: &[Eqn::T],
468    ) -> (
469        <Eqn::V as DefaultDenseMatrix>::M,
470        <Eqn::V as DefaultDenseMatrix>::M,
471    )
472    where
473        Eqn: OdeEquationsImplicitAdjoint + 'a,
474        LS: LinearSolver<Eqn::M>,
475        Eqn::M: DefaultSolver,
476        Eqn::V: DefaultDenseMatrix,
477        for<'b> &'b Eqn::V: VectorRef<Eqn::V>,
478        for<'b> &'b Eqn::M: MatrixRef<Eqn::M>,
479    {
480        let nparams = problem.eqn.nparams();
481        let nout = 2;
482        let ctx = problem.eqn.context();
483        let mut dgdp = <Eqn::V as DefaultDenseMatrix>::M::zeros(nparams, nout, ctx.clone());
484
485        let mut p_0 = ctx.vector_zeros(nparams);
486        problem.eqn.get_params(&mut p_0);
487        let h_base = Eqn::T::from_f64(1e-10).unwrap();
488        let mut h = Eqn::V::from_element(nparams, h_base, ctx.clone());
489        h.axpy(h_base, &p_0, Eqn::T::one());
490        let mut p_data = p_0.clone();
491        p_data.axpy(Eqn::T::from_f64(0.1).unwrap(), &p_0, Eqn::T::one());
492        let p_base = p_0.clone();
493
494        problem.eqn.set_params(&p_data);
495        let data = solve_dense_with_single_reset_root::<Eqn, _, _>(
496            |state| match state {
497                Some(state) => problem.bdf_solver(state),
498                None => problem.bdf::<LS>(),
499            },
500            times,
501        );
502
503        for i in 0..nparams {
504            p_0.set_index(i, p_base.get_index(i) + h.get_index(i));
505            problem.eqn.set_params(&p_0);
506            let g_pos = {
507                let v = solve_dense_with_single_reset_root::<Eqn, _, _>(
508                    |state| match state {
509                        Some(state) => problem.bdf_solver(state),
510                        None => problem.bdf::<LS>(),
511                    },
512                    times,
513                );
514                sum_squares(&v, &data)
515            };
516
517            p_0.set_index(i, p_base.get_index(i) - h.get_index(i));
518            problem.eqn.set_params(&p_0);
519            let g_neg = {
520                let v = solve_dense_with_single_reset_root::<Eqn, _, _>(
521                    |state| match state {
522                        Some(state) => problem.bdf_solver(state),
523                        None => problem.bdf::<LS>(),
524                    },
525                    times,
526                );
527                sum_squares(&v, &data)
528            };
529
530            p_0.set_index(i, p_base.get_index(i));
531
532            let delta = (g_pos - g_neg) / Scale(Eqn::T::from_f64(2.).unwrap() * h.get_index(i));
533            for j in 0..nout {
534                dgdp.set_index(i, j, delta.get_index(j));
535            }
536        }
537        problem.eqn.set_params(&p_base);
538        (dgdp, data)
539    }
540
541    pub fn test_adjoint_sum_squares<'a, Eqn, SolverF, SolverB>(
542        backwards_solver: SolverB,
543        dgdp_check: <Eqn::V as DefaultDenseMatrix>::M,
544        forwards_soln: <Eqn::V as DefaultDenseMatrix>::M,
545        data: <Eqn::V as DefaultDenseMatrix>::M,
546        times: &[Eqn::T],
547    ) where
548        SolverF: OdeSolverMethod<'a, Eqn>,
549        SolverB: AdjointOdeSolverMethod<'a, Eqn, SolverF>,
550        Eqn: OdeEquationsImplicitAdjoint + 'a,
551        Eqn::V: DefaultDenseMatrix,
552        Eqn::M: DefaultSolver,
553    {
554        let nparams = dgdp_check.nrows();
555        let dgdu = dsum_squaresdp(&forwards_soln, &data);
556
557        let atol = Eqn::V::from_element(
558            nparams,
559            Eqn::T::from_f64(1e-6).unwrap(),
560            data.context().clone(),
561        );
562        let rtol = Eqn::T::from_f64(1e-6).unwrap();
563        let (state, _) = backwards_solver
564            .solve_adjoint_backwards_pass(times, dgdu.iter().collect::<Vec<_>>().as_slice())
565            .unwrap();
566        let gs_adj = state.into_common().sg;
567        #[allow(clippy::needless_range_loop)]
568        for j in 0..dgdp_check.ncols() {
569            gs_adj[j].assert_eq_norm(
570                &dgdp_check.column(j).into_owned(),
571                &atol,
572                rtol,
573                Eqn::T::from_f64(260.).unwrap(),
574            );
575        }
576    }
577
578    pub fn test_adjoint<'a, Eqn, SolverF, SolverB>(
579        backwards_solver: SolverB,
580        dgdp_check: <Eqn::V as DefaultDenseMatrix>::M,
581        factor: Eqn::T,
582    ) where
583        SolverF: OdeSolverMethod<'a, Eqn>,
584        SolverB: AdjointOdeSolverMethod<'a, Eqn, SolverF>,
585        Eqn: OdeEquationsImplicitAdjoint + 'a,
586        Eqn::V: DefaultDenseMatrix,
587        Eqn::M: DefaultSolver,
588    {
589        let nout = backwards_solver.problem().eqn.nout();
590        let atol = Eqn::V::from_element(
591            nout,
592            Eqn::T::from_f64(1e-6).unwrap(),
593            dgdp_check.context().clone(),
594        );
595        let rtol = Eqn::T::from_f64(1e-6).unwrap();
596        let (state, _) = backwards_solver
597            .solve_adjoint_backwards_pass(&[], &[])
598            .unwrap();
599        let gs_adj = state.into_common().sg;
600        #[allow(clippy::needless_range_loop)]
601        for j in 0..dgdp_check.ncols() {
602            gs_adj[j].assert_eq_norm(&dgdp_check.column(j).into_owned(), &atol, rtol, factor);
603        }
604    }
605
606    pub struct TestEqnInit<M: Matrix> {
607        ctx: M::C,
608    }
609
610    impl<M: Matrix> Op for TestEqnInit<M> {
611        type T = M::T;
612        type V = M::V;
613        type M = M;
614        type C = M::C;
615
616        fn nout(&self) -> usize {
617            1
618        }
619        fn nparams(&self) -> usize {
620            1
621        }
622        fn nstates(&self) -> usize {
623            1
624        }
625        fn context(&self) -> &Self::C {
626            &self.ctx
627        }
628    }
629
630    impl<M: Matrix> ConstantOp for TestEqnInit<M> {
631        fn call_inplace(&self, _t: Self::T, y: &mut Self::V) {
632            y.fill(M::T::one());
633        }
634    }
635
636    impl<M: Matrix> ConstantOpSens for TestEqnInit<M> {
637        fn sens_mul_inplace(&self, _t: Self::T, _v: &Self::V, sens: &mut Self::V) {
638            sens.fill(M::T::zero());
639        }
640    }
641
642    pub struct TestEqnRhs<M: Matrix> {
643        ctx: M::C,
644    }
645
646    impl<M: Matrix> Op for TestEqnRhs<M> {
647        type T = M::T;
648        type V = M::V;
649        type M = M;
650        type C = M::C;
651
652        fn nout(&self) -> usize {
653            1
654        }
655        fn nparams(&self) -> usize {
656            1
657        }
658        fn nstates(&self) -> usize {
659            1
660        }
661        fn context(&self) -> &Self::C {
662            &self.ctx
663        }
664    }
665
666    impl<M: Matrix> NonLinearOp for TestEqnRhs<M> {
667        fn call_inplace(&self, _x: &Self::V, _t: Self::T, y: &mut Self::V) {
668            y.fill(M::T::zero());
669        }
670    }
671
672    impl<M: Matrix> NonLinearOpJacobian for TestEqnRhs<M> {
673        fn jac_mul_inplace(&self, _x: &Self::V, _t: Self::T, _v: &Self::V, y: &mut Self::V) {
674            y.fill(M::T::zero());
675        }
676    }
677
678    impl<M: Matrix> NonLinearOpSens for TestEqnRhs<M> {
679        fn sens_mul_inplace(&self, _x: &Self::V, _t: Self::T, _v: &Self::V, sens: &mut Self::V) {
680            sens.fill(M::T::zero());
681        }
682    }
683
684    pub struct TestEqnOut<M: Matrix> {
685        ctx: M::C,
686    }
687
688    impl<M: Matrix> Op for TestEqnOut<M> {
689        type T = M::T;
690        type V = M::V;
691        type M = M;
692        type C = M::C;
693
694        fn nout(&self) -> usize {
695            1
696        }
697        fn nparams(&self) -> usize {
698            1
699        }
700        fn nstates(&self) -> usize {
701            1
702        }
703        fn context(&self) -> &Self::C {
704            &self.ctx
705        }
706    }
707
708    impl<M: Matrix> NonLinearOp for TestEqnOut<M> {
709        fn call_inplace(&self, x: &Self::V, _t: Self::T, y: &mut Self::V) {
710            y.copy_from(x);
711        }
712    }
713
714    impl<M: Matrix> NonLinearOpJacobian for TestEqnOut<M> {
715        fn jac_mul_inplace(&self, _x: &Self::V, _t: Self::T, v: &Self::V, y: &mut Self::V) {
716            y.copy_from(v);
717        }
718    }
719
720    impl<M: Matrix> NonLinearOpSens for TestEqnOut<M> {
721        fn sens_mul_inplace(&self, _x: &Self::V, _t: Self::T, _v: &Self::V, sens: &mut Self::V) {
722            sens.fill(M::T::zero());
723        }
724    }
725
726    pub struct TestEqn<M: Matrix> {
727        rhs: Rc<TestEqnRhs<M>>,
728        init: Rc<TestEqnInit<M>>,
729        out: Rc<TestEqnOut<M>>,
730        ctx: M::C,
731    }
732
733    impl<M: Matrix> TestEqn<M> {
734        pub fn new() -> Self {
735            let ctx = M::C::default();
736            Self {
737                rhs: Rc::new(TestEqnRhs { ctx: ctx.clone() }),
738                init: Rc::new(TestEqnInit { ctx: ctx.clone() }),
739                out: Rc::new(TestEqnOut { ctx: ctx.clone() }),
740                ctx,
741            }
742        }
743    }
744
745    impl<M: Matrix> Op for TestEqn<M> {
746        type T = M::T;
747        type V = M::V;
748        type M = M;
749        type C = M::C;
750        fn nout(&self) -> usize {
751            1
752        }
753        fn nparams(&self) -> usize {
754            1
755        }
756        fn nstates(&self) -> usize {
757            1
758        }
759        fn statistics(&self) -> crate::op::OpStatistics {
760            OpStatistics::default()
761        }
762        fn context(&self) -> &Self::C {
763            &self.ctx
764        }
765    }
766
767    impl<'a, M: Matrix> OdeEquationsRef<'a> for TestEqn<M> {
768        type Rhs = &'a TestEqnRhs<M>;
769        type Mass = ParameterisedOp<'a, UnitCallable<M>>;
770        type Root = ParameterisedOp<'a, UnitCallable<M>>;
771        type Init = &'a TestEqnInit<M>;
772        type Out = &'a TestEqnOut<M>;
773        type Reset = ParameterisedOp<'a, UnitCallable<M>>;
774    }
775
776    impl<M: Matrix> OdeEquations for TestEqn<M> {
777        fn rhs(&self) -> &TestEqnRhs<M> {
778            &self.rhs
779        }
780
781        fn mass(&self) -> Option<<Self as OdeEquationsRef<'_>>::Mass> {
782            None
783        }
784
785        fn root(&self) -> Option<<Self as OdeEquationsRef<'_>>::Root> {
786            None
787        }
788
789        fn init(&self) -> &TestEqnInit<M> {
790            &self.init
791        }
792
793        fn out(&self) -> Option<<Self as OdeEquationsRef<'_>>::Out> {
794            Some(&self.out)
795        }
796        fn set_params(&mut self, _p: &Self::V) {
797            unimplemented!()
798        }
799        fn get_params(&self, _p: &mut Self::V) {
800            unimplemented!()
801        }
802    }
803
804    pub fn test_problem<M: Matrix>(integrate_out: bool) -> OdeSolverProblem<TestEqn<M>> {
805        let eqn = TestEqn::<M>::new();
806        let atol = eqn
807            .context()
808            .vector_from_element(1, M::T::from_f64(1e-6).unwrap());
809        OdeSolverProblem::new(
810            eqn,
811            M::T::from_f64(1e-6).unwrap(),
812            atol,
813            None,
814            None,
815            None,
816            None,
817            None,
818            None,
819            M::T::zero(),
820            M::T::one(),
821            integrate_out,
822            Default::default(),
823            Default::default(),
824        )
825        .unwrap()
826    }
827
828    pub fn test_interpolate<'a, M: Matrix, Method: OdeSolverMethod<'a, TestEqn<M>>>(mut s: Method) {
829        let state = s.checkpoint();
830        let integrating_sens = !s.state().s.is_empty();
831        let integrating_out = s.problem().integrate_out;
832        let t0 = state.as_ref().t;
833        let t1 = t0 + M::T::from_f64(1e6).unwrap();
834        s.interpolate(t0)
835            .unwrap()
836            .assert_eq_st(state.as_ref().y, M::T::from_f64(1e-9).unwrap());
837        assert!(s.interpolate(t1).is_err());
838        assert!(s.interpolate_out(t1).is_err());
839        if integrating_sens {
840            assert!(s.interpolate_sens(t1).is_err());
841        } else {
842            assert!(s.interpolate_sens(t0).is_ok());
843        }
844        s.step().unwrap();
845        let tmid = t0 + (s.state().t - t0) / M::T::from_f64(2.0).unwrap();
846        assert!(s.interpolate(s.state().t).is_ok());
847        assert!(s.interpolate(tmid).is_ok());
848        if integrating_out {
849            assert!(s.interpolate_out(s.state().t).is_ok());
850        } else {
851            assert!(s.interpolate_out(s.state().t).is_err());
852        }
853        assert!(s.interpolate_sens(s.state().t).is_ok());
854        assert!(s.interpolate(s.state().t + t1).is_err());
855        assert!(s.interpolate_out(s.state().t + t1).is_err());
856        if integrating_sens {
857            assert!(s.interpolate_sens(s.state().t + t1).is_err());
858        } else {
859            assert!(s.interpolate_sens(s.state().t + t1).is_ok());
860        }
861
862        let mut y_wrong_length = M::V::zeros(2, s.problem().context().clone());
863        assert!(s
864            .interpolate_inplace(s.state().t, &mut y_wrong_length)
865            .is_err());
866        let mut g_wrong_length = M::V::zeros(2, s.problem().context().clone());
867        assert!(s
868            .interpolate_out_inplace(s.state().t, &mut g_wrong_length)
869            .is_err());
870        let mut s_wrong_length = vec![
871            M::V::zeros(1, s.problem().context().clone()),
872            M::V::zeros(1, s.problem().context().clone()),
873        ];
874        assert!(s
875            .interpolate_sens_inplace(s.state().t, &mut s_wrong_length)
876            .is_err());
877        let mut s_wrong_vec_length = if integrating_sens {
878            vec![M::V::zeros(2, s.problem().context().clone())]
879        } else {
880            vec![]
881        };
882        if integrating_sens {
883            assert!(s
884                .interpolate_sens_inplace(s.state().t, &mut s_wrong_vec_length)
885                .is_err());
886        } else {
887            assert!(s
888                .interpolate_sens_inplace(s.state().t, &mut s_wrong_vec_length)
889                .is_ok());
890        }
891
892        s.state_mut().y.fill(M::T::from_f64(3.0).unwrap());
893        assert!(s.interpolate(s.state().t).is_ok());
894        if integrating_out {
895            assert!(s.interpolate_out(s.state().t).is_ok());
896        }
897        if integrating_sens {
898            assert!(s.interpolate_sens(s.state().t).is_ok());
899        }
900        assert!(s.interpolate(tmid).is_err());
901        assert!(s.interpolate_out(tmid).is_err());
902        if integrating_sens {
903            assert!(s.interpolate_sens(tmid).is_err());
904        } else {
905            assert!(s.interpolate_sens(tmid).is_ok());
906        }
907    }
908
909    pub fn test_interpolate_dy<'a, M: Matrix, Method: OdeSolverMethod<'a, TestEqn<M>>>(
910        mut s: Method,
911    ) {
912        // Error before first step: t is in the future
913        let t_future = s.state().t + M::T::from_f64(1e6).unwrap();
914        assert!(s.interpolate_dy(t_future).is_err());
915
916        let t0 = s.state().t;
917        s.step().unwrap();
918        let t1 = s.state().t;
919        let dt = t1 - t0;
920        let tmid = t0 + dt / M::T::from_f64(2.0).unwrap();
921
922        // Wrong vector length should return error
923        let mut dy_wrong = M::V::zeros(2, s.problem().context().clone());
924        assert!(s.interpolate_dy_inplace(t1, &mut dy_wrong).is_err());
925
926        // t after current time should return error
927        assert!(s.interpolate_dy(t1 + M::T::from_f64(1e6).unwrap()).is_err());
928
929        // interpolate_dy should be consistent with finite-difference of interpolate (step 1)
930        let eps = dt.abs() * M::T::from_f64(1e-5).unwrap();
931        let y_plus = s.interpolate(tmid + eps).unwrap();
932        let y_minus = s.interpolate(tmid - eps).unwrap();
933        let fd_dy = (y_plus - y_minus) * Scale(M::T::one() / (M::T::from_f64(2.0).unwrap() * eps));
934        let dy = s.interpolate_dy(tmid).unwrap();
935        dy.assert_eq_norm(
936            &fd_dy,
937            &s.problem().atol,
938            s.problem().rtol,
939            M::T::from_f64(1e3).unwrap(),
940        );
941
942        // take a second step and check consistency again
943        let t1 = s.state().t;
944        s.step().unwrap();
945        let t2 = s.state().t;
946        let dt2 = t2 - t1;
947        let tmid2 = t1 + dt2 / M::T::from_f64(2.0).unwrap();
948        let eps2 = dt2.abs() * M::T::from_f64(1e-5).unwrap();
949        let y_plus = s.interpolate(tmid2 + eps2).unwrap();
950        let y_minus = s.interpolate(tmid2 - eps2).unwrap();
951        let fd_dy2 =
952            (y_plus - y_minus) * Scale(M::T::one() / (M::T::from_f64(2.0).unwrap() * eps2));
953        let dy2 = s.interpolate_dy(tmid2).unwrap();
954        dy2.assert_eq_norm(
955            &fd_dy2,
956            &s.problem().atol,
957            s.problem().rtol,
958            M::T::from_f64(1e3).unwrap(),
959        );
960    }
961
962    pub fn test_config<'a, Eqn: OdeEquations + 'a, Method: OdeSolverMethod<'a, Eqn>>(
963        mut s: Method,
964    ) {
965        *s.config_mut().as_base_mut().minimum_timestep = Eqn::T::from_f64(1.0e8).unwrap();
966        assert_eq!(
967            *s.config().as_base_ref().minimum_timestep,
968            Eqn::T::from_f64(1.0e8).unwrap()
969        );
970        // force a step size reduction
971        *s.state_mut().h = Eqn::T::from_f64(0.1).unwrap();
972
973        let mut failed = false;
974        for _ in 0..10 {
975            if let Err(DiffsolError::OdeSolverError(OdeSolverError::StepSizeTooSmall { time: _ })) =
976                s.step()
977            {
978                failed = true;
979                break;
980            }
981        }
982        assert!(failed);
983    }
984
985    pub fn test_state_mut<'a, M: Matrix, Method: OdeSolverMethod<'a, TestEqn<M>>>(mut s: Method) {
986        let state = s.checkpoint();
987        let state2 = s.state();
988        state2
989            .y
990            .assert_eq_st(state.as_ref().y, M::T::from_f64(1e-9).unwrap());
991        s.state_mut()
992            .y
993            .set_index(0, M::T::from_f64(std::f64::consts::PI).unwrap());
994        assert_eq!(
995            s.state_mut().y.get_index(0),
996            M::T::from_f64(std::f64::consts::PI).unwrap()
997        );
998    }
999
1000    #[cfg(feature = "diffsl-cranelift")]
1001    pub fn test_ball_bounce_problem<M: crate::MatrixHost<T = f64>>(
1002    ) -> OdeSolverProblem<crate::DiffSl<M, crate::CraneliftJitModule>> {
1003        crate::OdeBuilder::<M>::new()
1004            .build_from_diffsl(
1005                "
1006            g { 9.81 } h { 10.0 }
1007            u_i {
1008                x = h,
1009                v = 0,
1010            }
1011            F_i {
1012                v,
1013                -g,
1014            }
1015            stop {
1016                x,
1017            }
1018        ",
1019            )
1020            .unwrap()
1021    }
1022
1023    #[cfg(feature = "diffsl-cranelift")]
1024    pub fn test_ball_bounce<'a, M, Method>(mut solver: Method) -> (Vec<f64>, Vec<f64>, Vec<f64>)
1025    where
1026        M: crate::MatrixHost<T = f64>,
1027        M: DefaultSolver<T = f64>,
1028        M::V: DefaultDenseMatrix<T = f64>,
1029        Method: OdeSolverMethod<'a, crate::DiffSl<M, crate::CraneliftJitModule>>,
1030    {
1031        let e = 0.8;
1032
1033        let final_time = 2.5;
1034
1035        // solve and apply the remaining doses
1036        solver.set_stop_time(final_time).unwrap();
1037        loop {
1038            match solver.step() {
1039                Ok(OdeSolverStopReason::InternalTimestep) => (),
1040                Ok(OdeSolverStopReason::RootFound(t, _)) => {
1041                    // get the state when the event occurred
1042                    let mut y = solver.interpolate(t).unwrap();
1043
1044                    // update the velocity of the ball
1045                    y.set_index(1, y.get_index(1) * -e);
1046
1047                    // make sure the ball is above the ground
1048                    y.set_index(0, y.get_index(0).max(f64::EPSILON));
1049
1050                    // set the state to the updated state
1051                    solver.state_mut().y.copy_from(&y);
1052                    solver.state_mut().dy.set_index(0, y.get_index(1));
1053                    *solver.state_mut().t = t;
1054
1055                    break;
1056                }
1057                Ok(OdeSolverStopReason::TstopReached) => break,
1058                Err(_) => panic!("unexpected solver error"),
1059            }
1060        }
1061        // do three more steps after the 1st bound and many sure they are correct
1062        let mut x = vec![];
1063        let mut v = vec![];
1064        let mut t = vec![];
1065        for _ in 0..3 {
1066            let ret = solver.step();
1067            x.push(solver.state().y.get_index(0));
1068            v.push(solver.state().y.get_index(1));
1069            t.push(solver.state().t);
1070            match ret {
1071                Ok(OdeSolverStopReason::InternalTimestep) => (),
1072                Ok(OdeSolverStopReason::RootFound(_, _)) => {
1073                    panic!("should be an internal timestep but found a root")
1074                }
1075                Ok(OdeSolverStopReason::TstopReached) => break,
1076                _ => panic!("should be an internal timestep"),
1077            }
1078        }
1079        (x, v, t)
1080    }
1081
1082    pub fn test_checkpointing<'a, M, Method, Eqn>(
1083        soln: OdeSolverSolution<M::V>,
1084        mut solver1: Method,
1085        mut solver2: Method,
1086    ) where
1087        M: Matrix + DefaultSolver,
1088        Method: OdeSolverMethod<'a, Eqn>,
1089        Eqn: OdeEquationsImplicit<M = M, T = M::T, V = M::V> + 'a,
1090    {
1091        let half_i = soln.solution_points.len() / 2;
1092        let half_t = soln.solution_points[half_i].t;
1093        while solver1.state().t <= half_t {
1094            solver1.step().unwrap();
1095        }
1096        let checkpoint = solver1.checkpoint();
1097        let checkpoint_t = checkpoint.as_ref().t;
1098        solver2.set_state(checkpoint);
1099
1100        // carry on solving with both solvers, they should produce about the same results (probably might diverge a bit, but should always match the solution)
1101        for point in soln.solution_points.iter().skip(half_i + 1) {
1102            // point should be past checkpoint
1103            if point.t < checkpoint_t {
1104                continue;
1105            }
1106            while solver2.state().t < point.t {
1107                solver1.step().unwrap();
1108                solver2.step().unwrap();
1109                let time_error = (solver1.state().t - solver2.state().t).abs()
1110                    / (solver1.state().t.abs() * solver1.problem().rtol
1111                        + solver1.problem().atol.get_index(0));
1112                assert!(
1113                    time_error < M::T::from_f64(20.0).unwrap(),
1114                    "time_error: {} at t = {}",
1115                    time_error,
1116                    solver1.state().t
1117                );
1118                solver1.state().y.assert_eq_norm(
1119                    solver2.state().y,
1120                    &solver1.problem().atol,
1121                    solver1.problem().rtol,
1122                    M::T::from_f64(20.0).unwrap(),
1123                );
1124            }
1125            let soln = solver1.interpolate(point.t).unwrap();
1126            soln.assert_eq_norm(
1127                &point.state,
1128                &solver1.problem().atol,
1129                solver1.problem().rtol,
1130                M::T::from_f64(15.0).unwrap(),
1131            );
1132            let soln = solver2.interpolate(point.t).unwrap();
1133            soln.assert_eq_norm(
1134                &point.state,
1135                &solver1.problem().atol,
1136                solver1.problem().rtol,
1137                M::T::from_f64(15.0).unwrap(),
1138            );
1139        }
1140    }
1141
1142    pub fn test_state_mut_on_problem<'a, Eqn, Method>(
1143        mut s: Method,
1144        soln: OdeSolverSolution<Eqn::V>,
1145    ) where
1146        Eqn: OdeEquationsImplicit + 'a,
1147        Eqn::M: DefaultSolver,
1148        Method: OdeSolverMethod<'a, Eqn>,
1149        Eqn::V: DefaultDenseMatrix,
1150    {
1151        // save state and solve for a little bit
1152        let state = s.checkpoint();
1153        s.solve(Eqn::T::one()).unwrap();
1154
1155        // reinit using state_mut
1156        s.state_mut().y.copy_from(state.as_ref().y);
1157        s.state_mut().dy.copy_from(state.as_ref().dy);
1158        *s.state_mut().t = state.as_ref().t;
1159
1160        // solve and check against solution
1161        for point in soln.solution_points.iter() {
1162            while s.state().t < point.t {
1163                s.step().unwrap();
1164            }
1165            let soln = s.interpolate(point.t).unwrap();
1166            let error = soln.clone() - &point.state;
1167            let error_norm = error
1168                .squared_norm(&error, &s.problem().atol, s.problem().rtol)
1169                .sqrt();
1170            assert!(
1171                error_norm < Eqn::T::from_f64(19.0).unwrap(),
1172                "error_norm: {} at t = {}",
1173                error_norm,
1174                point.t
1175            );
1176        }
1177    }
1178
1179    /// Test that `step()` returns `RootFound(t, index)` with the correct root index.
1180    ///
1181    /// The problem must have a root function with **two** outputs and **no** Reset:
1182    ///   - Root 0 fires first (at `t ≈ 5.108`, `y[0] ≈ 0.6` for the exponential-decay test model)
1183    ///   - Root 1 fires second
1184    ///
1185    /// The test asserts that the first `RootFound` reports index 0 and the time
1186    /// matches `t_root_0_expected` within `tol`.
1187    pub fn test_root_found_index<'a, Eqn, Method>(
1188        mut solver: Method,
1189        soln: &OdeSolverSolution<Eqn::V>,
1190        expected_root_index: usize,
1191        tol: Eqn::T,
1192    ) where
1193        Eqn: OdeEquations + 'a,
1194        Method: OdeSolverMethod<'a, Eqn>,
1195    {
1196        let t_root_expected = soln.solution_points[0].t;
1197        solver
1198            .set_stop_time(Eqn::T::from_f64(100.0).unwrap())
1199            .unwrap();
1200        loop {
1201            match solver.step().unwrap() {
1202                // RED: `RootFound` currently has one field; adding `index` makes this fail.
1203                OdeSolverStopReason::RootFound(t, index) => {
1204                    assert_eq!(
1205                        index, expected_root_index,
1206                        "expected root index {expected_root_index} but got {index}",
1207                    );
1208                    assert!(
1209                        (t - t_root_expected).abs() < tol,
1210                        "expected t ≈ {t_root_expected:?}, got {t:?}",
1211                    );
1212                    break;
1213                }
1214                OdeSolverStopReason::TstopReached => {
1215                    panic!("reached tstop without finding a root")
1216                }
1217                OdeSolverStopReason::InternalTimestep => {}
1218            }
1219        }
1220    }
1221
1222    /// Test that `solve()` automatically applies resets at roots and continues
1223    /// integrating until `final_time`.
1224    pub fn test_solve_with_reset<'a, Eqn, Method>(
1225        mut solver: Method,
1226        soln: &OdeSolverSolution<Eqn::V>,
1227        final_time: Eqn::T,
1228    ) where
1229        Eqn: OdeEquationsImplicit + 'a,
1230        Eqn::M: DefaultSolver,
1231        Eqn::V: DefaultDenseMatrix,
1232        Method: OdeSolverMethod<'a, Eqn>,
1233    {
1234        let (ys, ts, stop_reason) = solver.solve(final_time).unwrap();
1235        assert_eq!(stop_reason, OdeSolverStopReason::TstopReached);
1236        let t_last = *ts.last().unwrap();
1237        let time_tol = soln.rtol * final_time.abs() + soln.atol.get_index(0);
1238        assert!(
1239            (t_last - final_time).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1240            "expected solve() to reach final_time ≈ {:?}, got {:?}",
1241            final_time,
1242            t_last,
1243        );
1244        assert!(
1245            (solver.state().t - final_time).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1246            "expected solver state at final_time ≈ {:?}, got {:?}",
1247            final_time,
1248            solver.state().t,
1249        );
1250
1251        let expected = &soln.solution_points[0];
1252        let root_time_tol = soln.rtol * expected.t.abs() + soln.atol.get_index(0);
1253        let root_col = ts
1254            .iter()
1255            .position(|&t| (t - expected.t).abs() < Eqn::T::from_f64(30.0).unwrap() * root_time_tol)
1256            .expect("expected solve() output to include the second-root/reset time");
1257        let root_expected = Eqn::V::from_element(
1258            expected.state.len(),
1259            Eqn::T::from_f64(0.4).unwrap(),
1260            expected.state.context().clone(),
1261        );
1262        let root_state = ys.column(root_col).into_owned();
1263        let root_error = root_state - &root_expected;
1264        let root_error_norm = root_error
1265            .squared_norm(&root_expected, &soln.atol, soln.rtol)
1266            .sqrt();
1267        let error_threshold = Eqn::T::from_f64(20.0).unwrap();
1268        assert!(
1269            root_error_norm < error_threshold,
1270            "expected reset state y=0.4 at second-root time; WRMS error norm {root_error_norm:?} ≥ {error_threshold:?}",
1271        );
1272
1273        let reset_value = Eqn::T::from_f64(0.4).unwrap();
1274        let reset_tol = Eqn::T::from_f64(30.0).unwrap()
1275            * (soln.rtol * reset_value.abs() + soln.atol.get_index(0));
1276        let last_reset_col = (0..ts.len())
1277            .rev()
1278            .find(|&i| (ys.get_index(0, i) - reset_value).abs() < reset_tol)
1279            .expect("expected solve() output to include at least one reset state");
1280        let final_time_f64 = final_time.to_f64().unwrap();
1281        let last_reset_time_f64 = ts[last_reset_col].to_f64().unwrap();
1282        let expected_final_value =
1283            Eqn::T::from_f64(0.4 * (-0.1 * (final_time_f64 - last_reset_time_f64)).exp()).unwrap();
1284        let expected_final = Eqn::V::from_element(
1285            expected.state.len(),
1286            expected_final_value,
1287            expected.state.context().clone(),
1288        );
1289        let final_state = ys.column(ts.len() - 1).into_owned();
1290        let final_error = final_state - &expected_final;
1291        let final_error_norm = final_error
1292            .squared_norm(&expected_final, &soln.atol, soln.rtol)
1293            .sqrt();
1294        assert!(
1295            final_error_norm < error_threshold,
1296            "final state mismatch after automatic reset continuation: WRMS error norm {final_error_norm:?} ≥ {error_threshold:?}",
1297        );
1298    }
1299
1300    /// Test that `solve_dense()` automatically applies resets at roots and
1301    /// continues filling the requested evaluation times.
1302    pub fn test_solve_dense_with_reset<'a, Eqn, Method>(
1303        mut solver: Method,
1304        soln: &OdeSolverSolution<Eqn::V>,
1305    ) where
1306        Eqn: OdeEquationsImplicit + 'a,
1307        Eqn::M: DefaultSolver,
1308        Eqn::V: DefaultDenseMatrix,
1309        Method: OdeSolverMethod<'a, Eqn>,
1310    {
1311        let t_stop = soln.solution_points[0].t;
1312        let final_time = t_stop * Eqn::T::from_f64(2.0).unwrap();
1313        let mut probe_solver = solver.clone();
1314        let (probe_ys, probe_ts, probe_stop_reason) = probe_solver.solve(final_time).unwrap();
1315        assert_eq!(probe_stop_reason, OdeSolverStopReason::TstopReached);
1316
1317        let reset_time_tol =
1318            Eqn::T::from_f64(30.0).unwrap() * (soln.rtol * t_stop.abs() + soln.atol.get_index(0));
1319        let post_event_dt = Eqn::T::from_f64(1e-6).unwrap();
1320        let reset_value = Eqn::T::from_f64(0.4).unwrap();
1321        let reset_value_tol = Eqn::T::from_f64(30.0).unwrap()
1322            * (soln.rtol * reset_value.abs() + soln.atol.get_index(0));
1323        let reset_col = (0..probe_ts.len())
1324            .find(|&i| {
1325                (probe_ts[i] - t_stop).abs() < reset_time_tol
1326                    && (probe_ys.get_index(0, i) - reset_value).abs() < reset_value_tol
1327            })
1328            .expect("expected solve() probe output to contain the second-root reset state");
1329        let t_event = probe_ts[reset_col];
1330        let t_eval = vec![Eqn::T::zero(), t_event, t_event + post_event_dt, final_time];
1331
1332        let (ret, stop_reason) = solver.solve_dense(&t_eval).unwrap();
1333        assert_eq!(stop_reason, OdeSolverStopReason::TstopReached);
1334        assert!(
1335            ret.ncols() == t_eval.len(),
1336            "expected solve_dense() to fill all requested evaluation times"
1337        );
1338        let time_tol = soln.rtol * final_time.abs() + soln.atol.get_index(0);
1339        assert!(
1340            (solver.state().t - final_time).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1341            "expected solver state at final_time ≈ {:?}, got {:?}",
1342            final_time,
1343            solver.state().t,
1344        );
1345
1346        let error_threshold = Eqn::T::from_f64(20.0).unwrap();
1347        let pre_reset_state = ret.column(1).into_owned();
1348        let pre_reset_error = pre_reset_state - &soln.solution_points[0].state;
1349        let pre_reset_error_norm = pre_reset_error
1350            .squared_norm(&soln.solution_points[0].state, &soln.atol, soln.rtol)
1351            .sqrt();
1352        assert!(
1353            pre_reset_error_norm < error_threshold,
1354            "expected pre-reset state at event time; WRMS norm {pre_reset_error_norm:?} >= {error_threshold:?}",
1355        );
1356
1357        let expected_post_reset_value =
1358            reset_value * (-Eqn::T::from_f64(0.1).unwrap() * post_event_dt).exp();
1359        let expected_post_reset = Eqn::V::from_element(
1360            soln.solution_points[0].state.len(),
1361            expected_post_reset_value,
1362            soln.solution_points[0].state.context().clone(),
1363        );
1364        let post_reset_state = ret.column(2).into_owned();
1365        let post_reset_error = post_reset_state - &expected_post_reset;
1366        let post_reset_error_norm = post_reset_error
1367            .squared_norm(&expected_post_reset, &soln.atol, soln.rtol)
1368            .sqrt();
1369        assert!(
1370            post_reset_error_norm < error_threshold,
1371            "expected reset state just after event time; WRMS norm {post_reset_error_norm:?} >= {error_threshold:?}",
1372        );
1373    }
1374
1375    /// Test that `solve_dense_sensitivities()` applies root-aware resets and
1376    /// continues filling the requested evaluation times.
1377    pub fn test_solve_dense_sensitivities_with_reset<'a, Eqn, Method>(
1378        mut solver: Method,
1379        soln: &OdeSolverSolution<Eqn::V>,
1380    ) where
1381        Eqn: OdeEquationsImplicitSens + 'a,
1382        Eqn::V: DefaultDenseMatrix,
1383        Eqn::M: DefaultSolver,
1384        Method: SensitivitiesOdeSolverMethod<'a, Eqn>,
1385    {
1386        let t_stop = soln.solution_points[0].t;
1387        let t_event = Eqn::T::from_f64(10.0 * (5.0_f64 / 3.0_f64).ln()).unwrap();
1388
1389        let post_event_dt = Eqn::T::from_f64(1e-6).unwrap();
1390        let t_eval = vec![Eqn::T::zero(), t_event, t_event + post_event_dt, t_stop];
1391        let (ret, ret_sens, stop_reason) = solver.solve_dense_sensitivities(&t_eval).unwrap();
1392        assert_eq!(stop_reason, OdeSolverStopReason::TstopReached);
1393        assert_eq!(ret.ncols(), t_eval.len());
1394        for ret_sens_j in &ret_sens {
1395            assert_eq!(ret_sens_j.ncols(), t_eval.len());
1396        }
1397
1398        let error_threshold = Eqn::T::from_f64(100.0).unwrap();
1399        let ctx = soln.solution_points[0].state.context().clone();
1400        let nstates = soln.solution_points[0].state.len();
1401
1402        let post_reset_y = Eqn::T::from_f64(2.6).unwrap()
1403            * (-Eqn::T::from_f64(0.1).unwrap() * post_event_dt).exp();
1404        let post_reset_t = t_event + post_event_dt;
1405        let expected_post_reset = Eqn::V::from_element(nstates, post_reset_y, ctx.clone());
1406        let expected_post_reset_sk =
1407            Eqn::V::from_element(nstates, -post_reset_y * post_reset_t, ctx.clone());
1408        let expected_post_reset_sy0 = Eqn::V::from_element(nstates, post_reset_y, ctx);
1409
1410        let col = 2;
1411        let ey = ret.column(col).into_owned() - &expected_post_reset;
1412        let esk = ret_sens[0].column(col).into_owned() - &expected_post_reset_sk;
1413        let esy0 = ret_sens[1].column(col).into_owned() - &expected_post_reset_sy0;
1414        let norm = (ey.squared_norm(&expected_post_reset, &soln.atol, soln.rtol)
1415            + esk.squared_norm(&expected_post_reset_sk, &soln.atol, soln.rtol)
1416            + esy0.squared_norm(&expected_post_reset_sy0, &soln.atol, soln.rtol))
1417        .sqrt();
1418        assert!(
1419            norm < error_threshold,
1420            "dense sensitivity mismatch just after reset; combined WRMS {norm:?} >= {error_threshold:?}",
1421        );
1422    }
1423
1424    pub fn test_solve_adjoint_with_single_reset_root<
1425        'a,
1426        Eqn,
1427        MethodF,
1428        MethodB,
1429        BuildForward,
1430        BuildAdjointState,
1431        BuildAdjointFromState,
1432    >(
1433        build_forward: BuildForward,
1434        soln: &OdeSolverSolution<Eqn::V>,
1435        build_adjoint_state: BuildAdjointState,
1436        build_adjoint_from_state: BuildAdjointFromState,
1437        use_replay_solver: bool,
1438    ) where
1439        Eqn: OdeEquationsImplicitAdjoint + 'a,
1440        Eqn::M: DefaultSolver,
1441        Eqn::V: DefaultDenseMatrix,
1442        MethodF: OdeSolverMethod<'a, Eqn>,
1443        MethodB: AdjointOdeSolverMethod<'a, Eqn, MethodF, State = MethodF::State>,
1444        BuildForward: Fn(Option<MethodF::State>) -> Result<MethodF, DiffsolError>,
1445        BuildAdjointState:
1446            Fn(&mut AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodF::State, DiffsolError>,
1447        BuildAdjointFromState:
1448            Fn(MethodF::State, AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodB, DiffsolError>,
1449    {
1450        let expected_out = &soln.solution_points[0];
1451        let forward_stop_time = expected_out.t + Eqn::T::from_f64(1.0).unwrap();
1452
1453        let mut forward_solver = build_forward(None).unwrap();
1454        let (checkpointers, _forward_y, _forward_t, stop_reason) = forward_solver
1455            .solve_with_checkpointing(forward_stop_time, None)
1456            .unwrap();
1457        assert_eq!(stop_reason, OdeSolverStopReason::TstopReached);
1458        assert!(
1459            checkpointers.len() >= 3,
1460            "expected checkpointing path to include the two reset events"
1461        );
1462        let problem = forward_solver.problem();
1463        let post_reset_solver = forward_solver.clone();
1464        let post_reset_root_idx = checkpointers[1]
1465            .terminal_reset_root_idx()
1466            .expect("second reset segment should record its terminal root index");
1467        let final_forward_state = checkpointers[1].last_checkpoint().clone();
1468        let t_second_root = final_forward_state.as_ref().t;
1469
1470        let out_error = final_forward_state.as_ref().g.clone() - &expected_out.state;
1471        let out_norm = out_error
1472            .squared_norm(&expected_out.state, &soln.atol, soln.rtol)
1473            .sqrt();
1474        assert!(
1475            out_norm < Eqn::T::from_f64(50.0).unwrap(),
1476            "forward integrated output mismatch at second root: actual {:?}, expected {:?}, WRMS {out_norm:?}",
1477            final_forward_state.as_ref().g,
1478            expected_out.state,
1479        );
1480        let time_tol = soln.rtol * expected_out.t.abs() + soln.atol.get_index(0);
1481        assert!(
1482            (t_second_root - expected_out.t).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1483            "expected second root time ≈ {:?}, got {:?}",
1484            expected_out.t,
1485            t_second_root,
1486        );
1487
1488        let adjoint_checkpointers = checkpointers.into_iter().take(2).collect::<Vec<_>>();
1489
1490        // make a broken adjoint that is missing the reset root metadata on the first segment, which should cause an error
1491        let mut missing_metadata_checkpointers = adjoint_checkpointers.clone();
1492        missing_metadata_checkpointers[0].clear_terminal_reset_root_idx();
1493        let missing_metadata_solver = use_replay_solver.then(|| post_reset_solver.clone());
1494        let mut missing_metadata_adjoint_eqn = problem.adjoint_equations(
1495            missing_metadata_checkpointers,
1496            missing_metadata_solver,
1497            None,
1498        );
1499        let mut missing_metadata_adjoint_state =
1500            build_adjoint_state(&mut missing_metadata_adjoint_eqn).unwrap();
1501        missing_metadata_adjoint_state
1502            .as_mut()
1503            .state_mut_adjoint_terminal_root(
1504                &problem.eqn,
1505                post_reset_root_idx,
1506                &final_forward_state,
1507                problem.integrate_out,
1508            )
1509            .unwrap();
1510        let missing_metadata_adjoint =
1511            build_adjoint_from_state(missing_metadata_adjoint_state, missing_metadata_adjoint_eqn)
1512                .unwrap();
1513        let missing_metadata_err =
1514            match missing_metadata_adjoint.solve_adjoint_backwards_pass(&[], &[]) {
1515                Ok(_) => panic!("expected missing reset metadata error"),
1516                Err(err) => err,
1517            };
1518        assert!(
1519            format!("{missing_metadata_err:?}").contains("Missing reset root metadata"),
1520            "expected missing reset metadata error, got {missing_metadata_err:?}",
1521        );
1522
1523        // now build the correct adjoint and check that it produces the correct gradient
1524        let adjoint_solver = use_replay_solver.then_some(post_reset_solver);
1525        let mut adjoint_eqn =
1526            problem.adjoint_equations(adjoint_checkpointers, adjoint_solver, None);
1527        let mut adjoint_state = build_adjoint_state(&mut adjoint_eqn).unwrap();
1528        adjoint_state
1529            .as_mut()
1530            .state_mut_adjoint_terminal_root(
1531                &problem.eqn,
1532                post_reset_root_idx,
1533                &final_forward_state,
1534                problem.integrate_out,
1535            )
1536            .unwrap();
1537        let adjoint = build_adjoint_from_state(adjoint_state, adjoint_eqn).unwrap();
1538        let (adjoint_state, _) = adjoint.solve_adjoint_backwards_pass(&[], &[]).unwrap();
1539
1540        let t0 = problem.t0;
1541        let ctx = problem.context().clone();
1542
1543        let sens_points = soln.sens_solution_points.as_ref().unwrap();
1544        let expected_grad = Eqn::V::from_vec(
1545            sens_points
1546                .iter()
1547                .map(|pts| pts[0].state.get_index(0))
1548                .collect(),
1549            ctx.clone(),
1550        );
1551        let atol = Eqn::V::from_element(expected_grad.len(), Eqn::T::from_f64(1e-6).unwrap(), ctx);
1552        let t0_tol = Eqn::T::from_f64(10.0).unwrap() * Eqn::T::EPSILON;
1553        assert!(
1554            (adjoint_state.as_ref().t - t0).abs() <= t0_tol,
1555            "expected adjoint final time {:?}, got {:?}",
1556            t0,
1557            adjoint_state.as_ref().t,
1558        );
1559        adjoint_state.as_ref().sg[0].assert_eq_norm(
1560            &expected_grad,
1561            &atol,
1562            Eqn::T::from_f64(1e-6).unwrap(),
1563            Eqn::T::from_f64(60.0).unwrap(),
1564        );
1565    }
1566
1567    #[allow(clippy::too_many_arguments)]
1568    pub fn test_solve_adjoint_sum_squares_with_single_reset_root<
1569        'a,
1570        Eqn,
1571        MethodF,
1572        MethodB,
1573        BuildForward,
1574        BuildAdjointState,
1575        BuildAdjointFromState,
1576    >(
1577        build_forward: BuildForward,
1578        soln: &OdeSolverSolution<Eqn::V>,
1579        build_adjoint_state: BuildAdjointState,
1580        build_adjoint_from_state: BuildAdjointFromState,
1581        use_replay_solver: bool,
1582        dgdp_check: <Eqn::V as DefaultDenseMatrix>::M,
1583        data: <Eqn::V as DefaultDenseMatrix>::M,
1584        times: &[Eqn::T],
1585    ) where
1586        Eqn: OdeEquationsImplicitAdjoint + 'a,
1587        Eqn::M: DefaultSolver,
1588        Eqn::V: DefaultDenseMatrix,
1589        MethodF: OdeSolverMethod<'a, Eqn>,
1590        MethodB: AdjointOdeSolverMethod<'a, Eqn, MethodF, State = MethodF::State>,
1591        BuildForward: Fn(Option<MethodF::State>) -> Result<MethodF, DiffsolError>,
1592        BuildAdjointState:
1593            Fn(&mut AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodF::State, DiffsolError>,
1594        BuildAdjointFromState:
1595            Fn(MethodF::State, AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodB, DiffsolError>,
1596    {
1597        let expected_out = &soln.solution_points[0];
1598        let forward_stop_time = expected_out.t + Eqn::T::from_f64(1.0).unwrap();
1599        let forwards_soln =
1600            solve_dense_with_single_reset_root::<Eqn, MethodF, _>(&build_forward, times);
1601        assert_eq!(
1602            forwards_soln.ncols(),
1603            times.len(),
1604            "expected stitched forward samples to cover every requested observation time",
1605        );
1606        let dgdu = dsum_squaresdp(&forwards_soln, &data);
1607        let dgdu_refs = dgdu.iter().collect::<Vec<_>>();
1608
1609        let mut forward_solver = build_forward(None).unwrap();
1610        let (checkpointers, _forward_y, _forward_t, stop_reason) = forward_solver
1611            .solve_with_checkpointing(forward_stop_time, None)
1612            .unwrap();
1613        assert_eq!(stop_reason, OdeSolverStopReason::TstopReached);
1614        assert!(
1615            checkpointers.len() >= 3,
1616            "expected checkpointing path to include the two reset events"
1617        );
1618        let problem = forward_solver.problem();
1619        let post_reset_solver = forward_solver.clone();
1620        let post_reset_root_idx = checkpointers[1]
1621            .terminal_reset_root_idx()
1622            .expect("second reset segment should record its terminal root index");
1623        let final_forward_state = checkpointers[1].last_checkpoint().clone();
1624        let t_second_root = final_forward_state.as_ref().t;
1625
1626        let time_tol = soln.rtol * expected_out.t.abs() + soln.atol.get_index(0);
1627        assert!(
1628            (t_second_root - expected_out.t).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1629            "expected second root time ≈ {:?}, got {:?}",
1630            expected_out.t,
1631            t_second_root,
1632        );
1633
1634        let adjoint_solver = use_replay_solver.then_some(post_reset_solver);
1635        let mut adjoint_eqn = problem.adjoint_equations(
1636            checkpointers.into_iter().take(2).collect(),
1637            adjoint_solver,
1638            Some(dgdu.len()),
1639        );
1640        let mut adjoint_state = build_adjoint_state(&mut adjoint_eqn).unwrap();
1641        adjoint_state
1642            .as_mut()
1643            .state_mut_adjoint_terminal_root(
1644                &problem.eqn,
1645                post_reset_root_idx,
1646                &final_forward_state,
1647                problem.integrate_out,
1648            )
1649            .unwrap();
1650        let adjoint = build_adjoint_from_state(adjoint_state, adjoint_eqn).unwrap();
1651        let (adjoint_state, _) = adjoint
1652            .solve_adjoint_backwards_pass(times, dgdu_refs.as_slice())
1653            .unwrap();
1654
1655        let t0 = problem.t0;
1656        let ctx = problem.context().clone();
1657
1658        let nparams = dgdp_check.nrows();
1659        let atol = Eqn::V::from_element(nparams, Eqn::T::from_f64(1e-6).unwrap(), ctx);
1660        let t0_tol = Eqn::T::from_f64(10.0).unwrap() * Eqn::T::EPSILON;
1661        assert!(
1662            (adjoint_state.as_ref().t - t0).abs() <= t0_tol,
1663            "expected adjoint final time {:?}, got {:?}",
1664            t0,
1665            adjoint_state.as_ref().t,
1666        );
1667        #[allow(clippy::needless_range_loop)]
1668        for j in 0..dgdp_check.ncols() {
1669            adjoint_state.as_ref().sg[j].assert_eq_norm(
1670                &dgdp_check.column(j).into_owned(),
1671                &atol,
1672                Eqn::T::from_f64(1e-6).unwrap(),
1673                Eqn::T::from_f64(260.0).unwrap(),
1674            );
1675        }
1676    }
1677
1678    pub fn test_solve_soln_adjoint_with_single_reset_root<
1679        'a,
1680        Eqn,
1681        MethodF,
1682        MethodB,
1683        BuildForward,
1684        BuildAdjointState,
1685        BuildAdjointFromState,
1686    >(
1687        build_forward: BuildForward,
1688        soln: &OdeSolverSolution<Eqn::V>,
1689        build_adjoint_state: BuildAdjointState,
1690        build_adjoint_from_state: BuildAdjointFromState,
1691        use_replay_solver: bool,
1692    ) where
1693        Eqn: OdeEquationsImplicitAdjoint + 'a,
1694        Eqn::M: DefaultSolver,
1695        Eqn::V: DefaultDenseMatrix,
1696        MethodF: OdeSolverMethod<'a, Eqn>,
1697        MethodB: AdjointOdeSolverMethod<'a, Eqn, MethodF, State = MethodF::State>,
1698        BuildForward: Fn(Option<MethodF::State>) -> Result<MethodF, DiffsolError>,
1699        BuildAdjointState:
1700            Fn(&mut AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodF::State, DiffsolError>,
1701        BuildAdjointFromState:
1702            Fn(MethodF::State, AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodB, DiffsolError>,
1703    {
1704        let expected_out = &soln.solution_points[0];
1705        let forward_stop_time = expected_out.t + Eqn::T::from_f64(1.0).unwrap();
1706        let mut forward_soln = Solution::<Eqn::V>::new(forward_stop_time);
1707        let mut checkpointers = Vec::new();
1708
1709        let first_forward_solver = build_forward(None)
1710            .unwrap()
1711            .solve_soln_with_checkpointing(&mut forward_soln, &mut checkpointers, None)
1712            .unwrap();
1713        let first_root_idx = match forward_soln.stop_reason {
1714            Some(OdeSolverStopReason::RootFound(_, idx)) => idx,
1715            Some(reason) => {
1716                panic!("expected first staged solve to stop at reset root, got {reason:?}")
1717            }
1718            None => panic!("first staged solve did not set a stop reason"),
1719        };
1720        assert_eq!(checkpointers.len(), 1);
1721        assert_eq!(
1722            checkpointers[0].terminal_reset_root_idx(),
1723            Some(first_root_idx)
1724        );
1725
1726        let state_after_reset = state_after_manual_reset::<Eqn, MethodF>(&first_forward_solver);
1727        let terminal_forward_solver = build_forward(Some(state_after_reset))
1728            .unwrap()
1729            .solve_soln_with_checkpointing(&mut forward_soln, &mut checkpointers, None)
1730            .unwrap();
1731        let terminal_root_idx = match forward_soln.stop_reason {
1732            Some(OdeSolverStopReason::RootFound(_, idx)) => idx,
1733            Some(reason) => {
1734                panic!("expected second staged solve to stop at terminal root, got {reason:?}")
1735            }
1736            None => panic!("second staged solve did not set a stop reason"),
1737        };
1738        assert_eq!(checkpointers.len(), 2);
1739        assert_eq!(
1740            checkpointers[1].terminal_reset_root_idx(),
1741            Some(terminal_root_idx)
1742        );
1743
1744        let problem = terminal_forward_solver.problem();
1745        let final_forward_state = terminal_forward_solver.state_clone();
1746        let t_second_root = final_forward_state.as_ref().t;
1747        let out_error = final_forward_state.as_ref().g.clone() - &expected_out.state;
1748        let out_norm = out_error
1749            .squared_norm(&expected_out.state, &soln.atol, soln.rtol)
1750            .sqrt();
1751        assert!(
1752            out_norm < Eqn::T::from_f64(50.0).unwrap(),
1753            "forward integrated output mismatch at terminal root: actual {:?}, expected {:?}, WRMS {out_norm:?}",
1754            final_forward_state.as_ref().g,
1755            expected_out.state,
1756        );
1757        let time_tol = soln.rtol * expected_out.t.abs() + soln.atol.get_index(0);
1758        assert!(
1759            (t_second_root - expected_out.t).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1760            "expected terminal root time ≈ {:?}, got {:?}",
1761            expected_out.t,
1762            t_second_root,
1763        );
1764
1765        let adjoint_solver = use_replay_solver.then_some(terminal_forward_solver.clone());
1766        let mut adjoint_eqn = problem.adjoint_equations(checkpointers, adjoint_solver, None);
1767        let mut adjoint_state = build_adjoint_state(&mut adjoint_eqn).unwrap();
1768        adjoint_state
1769            .as_mut()
1770            .state_mut_adjoint_terminal_root(
1771                &problem.eqn,
1772                terminal_root_idx,
1773                &final_forward_state,
1774                problem.integrate_out,
1775            )
1776            .unwrap();
1777        let adjoint = build_adjoint_from_state(adjoint_state, adjoint_eqn).unwrap();
1778        let (adjoint_state, _) = adjoint.solve_adjoint_backwards_pass(&[], &[]).unwrap();
1779
1780        let t0 = problem.t0;
1781        let ctx = problem.context().clone();
1782        let sens_points = soln.sens_solution_points.as_ref().unwrap();
1783        let expected_grad = Eqn::V::from_vec(
1784            sens_points
1785                .iter()
1786                .map(|pts| pts[0].state.get_index(0))
1787                .collect(),
1788            ctx.clone(),
1789        );
1790        let atol = Eqn::V::from_element(expected_grad.len(), Eqn::T::from_f64(1e-6).unwrap(), ctx);
1791        let t0_tol = Eqn::T::from_f64(10.0).unwrap() * Eqn::T::EPSILON;
1792        assert!(
1793            (adjoint_state.as_ref().t - t0).abs() <= t0_tol,
1794            "expected adjoint final time {:?}, got {:?}",
1795            t0,
1796            adjoint_state.as_ref().t,
1797        );
1798        adjoint_state.as_ref().sg[0].assert_eq_norm(
1799            &expected_grad,
1800            &atol,
1801            Eqn::T::from_f64(1e-6).unwrap(),
1802            Eqn::T::from_f64(60.0).unwrap(),
1803        );
1804    }
1805
1806    #[allow(clippy::too_many_arguments)]
1807    pub fn test_solve_soln_adjoint_sum_squares_with_single_reset_root<
1808        'a,
1809        Eqn,
1810        MethodF,
1811        MethodB,
1812        BuildForward,
1813        BuildAdjointState,
1814        BuildAdjointFromState,
1815    >(
1816        build_forward: BuildForward,
1817        soln: &OdeSolverSolution<Eqn::V>,
1818        build_adjoint_state: BuildAdjointState,
1819        build_adjoint_from_state: BuildAdjointFromState,
1820        use_replay_solver: bool,
1821        dgdp_check: <Eqn::V as DefaultDenseMatrix>::M,
1822        data: <Eqn::V as DefaultDenseMatrix>::M,
1823        times: &[Eqn::T],
1824    ) where
1825        Eqn: OdeEquationsImplicitAdjoint + 'a,
1826        Eqn::M: DefaultSolver,
1827        Eqn::V: DefaultDenseMatrix,
1828        MethodF: OdeSolverMethod<'a, Eqn>,
1829        MethodB: AdjointOdeSolverMethod<'a, Eqn, MethodF, State = MethodF::State>,
1830        BuildForward: Fn(Option<MethodF::State>) -> Result<MethodF, DiffsolError>,
1831        BuildAdjointState:
1832            Fn(&mut AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodF::State, DiffsolError>,
1833        BuildAdjointFromState:
1834            Fn(MethodF::State, AdjointEquations<'a, Eqn, MethodF>) -> Result<MethodB, DiffsolError>,
1835    {
1836        let expected_out = &soln.solution_points[0];
1837        let forward_stop_time = expected_out.t + Eqn::T::from_f64(1.0).unwrap();
1838        let mut forward_soln = Solution::<Eqn::V>::new_dense(times.to_vec()).unwrap();
1839        let mut checkpointers = Vec::new();
1840
1841        let first_forward_solver = build_forward(None)
1842            .unwrap()
1843            .solve_soln_with_checkpointing(&mut forward_soln, &mut checkpointers, None)
1844            .unwrap();
1845        let first_root_idx = match forward_soln.stop_reason {
1846            Some(OdeSolverStopReason::RootFound(_, idx)) => idx,
1847            Some(reason) => {
1848                panic!("expected first staged solve to stop at reset root, got {reason:?}")
1849            }
1850            None => panic!("first staged solve did not set a stop reason"),
1851        };
1852        assert_eq!(checkpointers.len(), 1);
1853        assert_eq!(
1854            checkpointers[0].terminal_reset_root_idx(),
1855            Some(first_root_idx)
1856        );
1857
1858        let state_after_reset = state_after_manual_reset::<Eqn, MethodF>(&first_forward_solver);
1859        build_forward(Some(state_after_reset.clone()))
1860            .unwrap()
1861            .solve_soln(&mut forward_soln)
1862            .unwrap();
1863        assert!(forward_soln.is_complete());
1864        assert_eq!(
1865            forward_soln.stop_reason,
1866            Some(OdeSolverStopReason::TstopReached)
1867        );
1868
1869        let mut terminal_soln = Solution::<Eqn::V>::new(forward_stop_time);
1870        let terminal_forward_solver = build_forward(Some(state_after_reset))
1871            .unwrap()
1872            .solve_soln_with_checkpointing(&mut terminal_soln, &mut checkpointers, None)
1873            .unwrap();
1874        let terminal_root_idx = match terminal_soln.stop_reason {
1875            Some(OdeSolverStopReason::RootFound(_, idx)) => idx,
1876            Some(reason) => {
1877                panic!("expected terminal staged solve to stop at root, got {reason:?}")
1878            }
1879            None => panic!("terminal staged solve did not set a stop reason"),
1880        };
1881        assert_eq!(checkpointers.len(), 2);
1882        assert_eq!(
1883            checkpointers.last().unwrap().terminal_reset_root_idx(),
1884            Some(terminal_root_idx)
1885        );
1886
1887        let dgdu_eval = dsum_squaresdp(&forward_soln.ys, &data);
1888        let dgdu_eval_refs = dgdu_eval.iter().collect::<Vec<_>>();
1889        let problem = terminal_forward_solver.problem();
1890        let final_forward_state = terminal_forward_solver.state_clone();
1891        let t_second_root = final_forward_state.as_ref().t;
1892        let time_tol = soln.rtol * expected_out.t.abs() + soln.atol.get_index(0);
1893        assert!(
1894            (t_second_root - expected_out.t).abs() < Eqn::T::from_f64(30.0).unwrap() * time_tol,
1895            "expected terminal root time ≈ {:?}, got {:?}",
1896            expected_out.t,
1897            t_second_root,
1898        );
1899
1900        let adjoint_solver = use_replay_solver.then_some(terminal_forward_solver.clone());
1901        let mut adjoint_eqn =
1902            problem.adjoint_equations(checkpointers, adjoint_solver, Some(dgdu_eval_refs.len()));
1903        let mut adjoint_state = build_adjoint_state(&mut adjoint_eqn).unwrap();
1904        adjoint_state
1905            .as_mut()
1906            .state_mut_adjoint_terminal_root(
1907                &problem.eqn,
1908                terminal_root_idx,
1909                &final_forward_state,
1910                problem.integrate_out,
1911            )
1912            .unwrap();
1913        let adjoint = build_adjoint_from_state(adjoint_state, adjoint_eqn).unwrap();
1914        let (adjoint_state, _) = adjoint
1915            .solve_adjoint_backwards_pass(times, dgdu_eval_refs.as_slice())
1916            .unwrap();
1917
1918        let t0 = problem.t0;
1919        let ctx = problem.context().clone();
1920        let nparams = dgdp_check.nrows();
1921        let atol = Eqn::V::from_element(nparams, Eqn::T::from_f64(1e-6).unwrap(), ctx);
1922        let t0_tol = Eqn::T::from_f64(10.0).unwrap() * Eqn::T::EPSILON;
1923        assert!(
1924            (adjoint_state.as_ref().t - t0).abs() <= t0_tol,
1925            "expected adjoint final time {:?}, got {:?}",
1926            t0,
1927            adjoint_state.as_ref().t,
1928        );
1929        #[allow(clippy::needless_range_loop)]
1930        for j in 0..dgdp_check.ncols() {
1931            adjoint_state.as_ref().sg[j].assert_eq_norm(
1932                &dgdp_check.column(j).into_owned(),
1933                &atol,
1934                Eqn::T::from_f64(1e-6).unwrap(),
1935                Eqn::T::from_f64(260.0).unwrap(),
1936            );
1937        }
1938    }
1939}