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#[derive(Clone, Debug, Serialize, Default)]
28pub struct OdeSolverStatistics {
29 pub number_of_linear_solver_setups: usize,
31 pub number_of_steps: usize,
33 pub number_of_error_test_failures: usize,
35 pub number_of_nonlinear_solver_iterations: usize,
37 pub number_of_nonlinear_solver_fails: usize,
39 pub number_of_linear_solver_setups_from_checkpoint: usize,
41 pub number_of_linear_solver_setups_from_first_convergence_fail: usize,
43 pub number_of_linear_solver_setups_from_second_convergence_fail: usize,
45 pub number_of_linear_solver_setups_from_error_test_fail: usize,
47 pub number_of_linear_solver_setups_from_step_success: usize,
49}
50
51impl OdeSolverStatistics {
52 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 (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 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 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 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 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 assert!(s.interpolate_dy(t1 + M::T::from_f64(1e6).unwrap()).is_err());
928
929 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 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 *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 solver.set_stop_time(final_time).unwrap();
1037 loop {
1038 match solver.step() {
1039 Ok(OdeSolverStopReason::InternalTimestep) => (),
1040 Ok(OdeSolverStopReason::RootFound(t, _)) => {
1041 let mut y = solver.interpolate(t).unwrap();
1043
1044 y.set_index(1, y.get_index(1) * -e);
1046
1047 y.set_index(0, y.get_index(0).max(f64::EPSILON));
1049
1050 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 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 for point in soln.solution_points.iter().skip(half_i + 1) {
1102 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 let state = s.checkpoint();
1153 s.solve(Eqn::T::one()).unwrap();
1154
1155 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 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 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 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 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 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 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 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 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}