1use anyhow::anyhow;
2use std::fmt;
3use std::sync::Arc;
4use tch::IndexOp;
5use tch::Tensor;
6
7pub mod optimizers;
8
9#[cfg(feature = "warnings")]
10use tracing::warn;
11
12#[cfg(not(feature = "warnings"))]
13macro_rules! warn {
14 ($($arg:tt)*) => {};
15}
16
17#[cfg(test)]
18mod tests;
19
20fn validate_finite_tensor(tensor: &Tensor, context: &str) -> anyhow::Result<()> {
23 if tensor.isfinite().f_all()?.f_int64_value(&[])? == 0 {
24 anyhow::bail!(
25 "Non-finite values (NaN/Inf) detected in {}: tensor shape {:?}",
26 context,
27 tensor.size()
28 );
29 }
30 Ok(())
31}
32
33fn validate_finite_scalar(value: f64, context: &str) -> anyhow::Result<()> {
35 if !value.is_finite() {
36 anyhow::bail!("Non-finite value ({}) detected in {}", value, context);
37 }
38 Ok(())
39}
40
41pub enum Solver {
42 Euler {
43 step: f64,
44 },
45 RK4 {
46 step: f64,
47 },
48 ImplicitEuler {
49 step: f64,
50 optimizer: Arc<dyn optimizers::Optimizer>,
51 },
52 GLRK4 {
53 step: f64,
54 optimizer: Arc<dyn optimizers::Optimizer>,
55 },
56 RKF45 {
57 rtol: f64,
58 atol: f64,
59 min_step: f64,
60 safety_factor: f64,
61 },
62 ROW1 {
63 step: f64,
64 },
65}
66
67impl Solver {
68 pub fn solve(
69 &self,
70 f: tch::CModule,
71 x_span: (f64, f64),
72 y0: Tensor,
73 ) -> anyhow::Result<(Tensor, Tensor)> {
74 let kind = y0.kind();
75 let device = y0.device();
76
77 if !x_span.0.is_finite() || !x_span.1.is_finite() {
79 return Err(anyhow!("x_span must consist of finite values"));
80 }
81 if x_span.0 > x_span.1 {
82 return Err(anyhow!("x_span is not a valid interval"));
83 }
84
85 match self {
87 Self::Euler { step }
88 | Self::RK4 { step }
89 | Self::ImplicitEuler { step, .. }
90 | Self::GLRK4 { step, .. }
91 | Self::ROW1 { step } => {
92 if !step.is_finite() || *step <= 0.0 {
93 return Err(anyhow!(
94 "Step size must be a finite positive value, got {}",
95 step
96 ));
97 }
98 }
99
100 Self::RKF45 {
101 rtol,
102 atol,
103 min_step,
104 safety_factor,
105 } => {
106 if !rtol.is_finite() || *rtol <= 0.0 {
107 return Err(anyhow!(
108 "rtol must be a finite positive value, got {}",
109 rtol
110 ));
111 }
112
113 if !atol.is_finite() || *atol <= 0.0 {
114 return Err(anyhow!(
115 "atol must be a finite positive value, got {}",
116 atol
117 ));
118 }
119
120 if !min_step.is_finite() || *min_step <= 0.0 {
121 return Err(anyhow!(
122 "min_step must be a finite positive value, got {}",
123 min_step
124 ));
125 }
126
127 if !safety_factor.is_finite() || *safety_factor <= 0.0 {
128 return Err(anyhow!(
129 "safety_factor must be a finite positive value, got {}",
130 safety_factor
131 ));
132 }
133 }
134 }
135
136 validate_finite_tensor(&y0, "initial state y0")?;
138
139 let y0_size = y0.size();
140
141 if y0_size.len() != 1 {
142 return Err(anyhow!(
143 "y0 must be a one-dimensional tensor but it has {} dimensions",
144 y0_size.len()
145 ));
146 }
147
148 if kind != tch::Kind::Double
149 && kind != tch::Kind::Float
150 && kind != tch::Kind::BFloat16
151 && kind != tch::Kind::Half
152 {
153 return Err(anyhow!("y0 is of unsupported kind {:?}", y0.kind()));
154 }
155
156 let dy = f.forward_ts(&[
158 Tensor::from(x_span.0).to_kind(kind).to_device(device),
159 y0.copy(),
160 ])?;
161
162 let dy_size = dy.size();
163
164 if dy_size.len() != 1 {
165 return Err(anyhow!(
166 "Function `f` returns tensor of rank {}, expected one-dimensional tensor",
167 dy_size.len()
168 ));
169 }
170
171 if dy_size[0] != y0_size[0] {
172 return Err(anyhow!(
173 "Function `f` returns vector of length {}, expected vector of length {} (same as y0)",
174 dy_size[0],
175 y0_size[0]
176 ));
177 }
178
179 if dy.device() != device {
180 return Err(anyhow!(
181 "Function `f` returns tensor on device {:?}, expected tensor to be on device {:?} (same as y0)",
182 dy.device(),
183 device
184 ));
185 }
186
187 if dy.kind() != kind {
188 return Err(anyhow!(
189 "Function `f` returns tensor of kind {:?}, expected tensor to be of kind {:?} (same as y0)",
190 dy.kind(),
191 kind
192 ));
193 }
194
195 validate_finite_tensor(&dy, "derivative function output at initial point")?;
197
198 match self {
199 Self::Euler { step } => solve_euler(f, x_span, y0, *step),
200
201 Self::RK4 { step } => solve_rk4(f, x_span, y0, *step),
202
203 Self::ImplicitEuler { step, optimizer } => {
204 solve_implicit_euler(f, x_span, y0, *step, optimizer.as_ref())
205 }
206
207 Self::GLRK4 { step, optimizer } => {
208 solve_glrk4(f, x_span, y0, *step, optimizer.as_ref())
209 }
210
211 Self::RKF45 {
212 rtol,
213 atol,
214 min_step,
215 safety_factor,
216 } => solve_rkf45(f, x_span, y0, *rtol, *atol, *min_step, *safety_factor),
217
218 Self::ROW1 { step } => solve_row1(f, x_span, y0, *step),
219 }
220 }
221
222 pub fn stability_function(&self, x: f64) -> anyhow::Result<f64> {
223 if x > 0. {
224 anyhow::bail!("Stability function is not defined for positive numbers.");
225 }
226
227 Ok(match self {
228 Self::Euler { .. } => 1. + x,
229 Self::RK4 { .. } => 1. + (1. + (1. / 2. + (1. / 6. + (1. / 24.) * x) * x) * x) * x,
230 Self::ImplicitEuler { .. } => 1. / (1. - x),
231 Self::GLRK4 { .. } => (1. + x / 2. + x * x / 12.) / (1. - x / 2. + x * x / 12.),
232 Self::RKF45 { .. } => {
233 1. + (1.
234 + (1. / 2.
235 + (1. / 6. + (1. / 24. + (1. / 120. + (1. / 2080.) * x) * x) * x) * x)
236 * x)
237 * x
238 }
239 Self::ROW1 { .. } => 1. / (1. - x),
240 })
241 }
242
243 pub fn stability_constant(&self) -> f64 {
244 match self {
245 Self::Euler { .. } => 2f64,
246 Self::RK4 { .. } => 2.785293563f64,
247 Self::ImplicitEuler { .. } => f64::INFINITY,
248 Self::GLRK4 { .. } => f64::INFINITY,
249 Self::RKF45 { .. } => 3.677706621f64,
250 Self::ROW1 { .. } => f64::INFINITY,
251 }
252 }
253}
254
255impl fmt::Display for Solver {
256 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
257 match self {
258 Solver::Euler { step } => write!(f, "Euler(step={})", step),
259 Solver::RK4 { step } => write!(f, "RK4(step={})", step),
260 Solver::ImplicitEuler { step, optimizer } => {
261 write!(f, "ImplicitEuler(step={}, optimizer={})", step, optimizer)
262 }
263 Solver::GLRK4 { step, optimizer } => {
264 write!(f, "GLRK4(step={}, optimizer={})", step, optimizer)
265 }
266 Solver::RKF45 {
267 rtol,
268 atol,
269 min_step,
270 safety_factor,
271 } => write!(
272 f,
273 "RKF45(rtol={}, atol={}, min_step={}, safety_factor={})",
274 rtol, atol, min_step, safety_factor
275 ),
276 Solver::ROW1 { step } => write!(f, "ROW1(step={})", step),
277 }
278 }
279}
280
281fn solve_euler(
283 f: tch::CModule,
284 x_span: (f64, f64),
285 y0: Tensor,
286 step: f64,
287) -> anyhow::Result<(Tensor, Tensor)> {
288 let device = y0.device();
289 let kind = y0.kind();
290
291 let x_start = x_span.0;
292 let x_end = x_span.1;
293
294 let mut x = x_start;
295 let mut y = y0.copy();
296
297 let mut all_x = vec![x];
298 let mut all_y = vec![y.copy()];
299
300 let mut current_step = step;
301 let mut step_count: u64 = 0;
302
303 let mut warned_large_norm = false;
304 let mut warned_many_steps = false;
305
306 while x < x_end {
307 let remaining = x_end - x;
308 if remaining < current_step {
309 current_step = remaining;
310 }
311
312 let dy = f.forward_ts(&[Tensor::from(x).to_kind(kind).to_device(device), y.copy()])?;
313
314 validate_finite_tensor(&dy, "derivative from f(x, y) in Euler step")?;
315
316 let dy_size = dy.size();
317 let dy_rank = dy_size.len();
318 if dy_rank != 1 {
319 anyhow::bail!(
320 "Derivative CModule returned tensor of bad rank {}.",
321 dy_rank
322 );
323 }
324 if dy_size[0] != y0.size()[0] {
325 anyhow::bail!(
326 "Derivative CModule returned vector of bad length {}.",
327 dy_size[0]
328 );
329 }
330
331 y = &y + current_step * &dy;
333
334 validate_finite_tensor(&y, "state after Euler update (NaN/Inf propagating)")?;
336
337 x = &x + current_step;
338
339 let x_tensor = Tensor::from(x).to_kind(kind).to_device(device);
341 validate_finite_tensor(&x_tensor, "integration variable x in Euler step")?;
342
343 all_x.push(x);
344 all_y.push(y.copy());
345
346 step_count += 1;
347
348 let y_norm = y.f_norm()?.f_double_value(&[])?;
349
350 if !warned_large_norm && y_norm > 1e10 {
351 warn!(
352 "Euler: solution norm exceeded {:.1e} at x={:.3e}; the solution may be diverging.",
353 1e10, x
354 );
355 warned_large_norm = true;
356 }
357
358 if !warned_many_steps && step_count >= 100_000 {
359 warn!(
360 "Euler: reached {} steps; consider increasing step size or switching to a higher-order solver",
361 step_count
362 );
363 warned_many_steps = true;
364 }
365 }
366
367 Ok((
368 Tensor::f_from_slice(&all_x)?
369 .to_kind(kind)
370 .to_device(device),
371 Tensor::f_stack(&all_y, 0)?,
372 ))
373}
374
375fn solve_rk4(
377 f: tch::CModule,
378 x_span: (f64, f64),
379 y0: Tensor,
380 step: f64,
381) -> anyhow::Result<(Tensor, Tensor)> {
382 let device = y0.device();
383 let kind = y0.kind();
384
385 let x_start = x_span.0;
386 let x_end = x_span.1;
387
388 let mut x = x_start;
389 let mut y = y0.copy();
390
391 let mut all_x = vec![x];
392 let mut all_y = vec![y.copy()];
393
394 let mut current_step = step;
395 let mut step_count: u64 = 0;
396
397 let mut warned_large_norm = false;
398 let mut warned_many_steps = false;
399
400 while x < x_end {
401 let remaining = x_end - x;
402 if remaining < current_step {
403 current_step = remaining;
404 }
405
406 let k1 = f.forward_ts(&[Tensor::from(x).to_kind(kind).to_device(device), y.copy()])?;
408 validate_finite_tensor(&k1, "RK4 stage k1")?;
409
410 let k1_size = k1.size();
411 if k1_size.len() != 1 || k1_size[0] != y0.size()[0] {
412 anyhow::bail!("Derivative CModule returned tensor of wrong shape at stage k1");
413 }
414
415 let x_half = x + 0.5 * current_step;
417 let y_half: Tensor = &y + 0.5 * current_step * &k1;
418 validate_finite_tensor(&y_half, "RK4 intermediate state for k2")?;
419
420 let k2 = f.forward_ts(&[Tensor::from(x_half).to_kind(kind).to_device(device), y_half])?;
421 validate_finite_tensor(&k2, "RK4 stage k2")?;
422
423 let k2_size = k2.size();
424 if k2_size.len() != 1 || k2_size[0] != y0.size()[0] {
425 anyhow::bail!("Derivative CModule returned tensor of wrong shape at stage k2");
426 }
427
428 let x_half_again = x + 0.5 * current_step;
430 let y_half_again: Tensor = &y + 0.5 * current_step * &k2;
431 validate_finite_tensor(&y_half_again, "RK4 intermediate state for k3")?;
432
433 let k3 = f.forward_ts(&[
434 Tensor::from(x_half_again).to_kind(kind).to_device(device),
435 y_half_again,
436 ])?;
437 validate_finite_tensor(&k3, "RK4 stage k3")?;
438
439 let k3_size = k3.size();
440 if k3_size.len() != 1 || k3_size[0] != y0.size()[0] {
441 anyhow::bail!("Derivative CModule returned tensor of wrong shape at stage k3");
442 }
443
444 let x_full = x + current_step;
446 let y_full = &y + current_step * &k3;
447 validate_finite_tensor(&y_full, "RK4 intermediate state for k4")?;
448
449 let k4 = f.forward_ts(&[Tensor::from(x_full).to_kind(kind).to_device(device), y_full])?;
450 validate_finite_tensor(&k4, "RK4 stage k4")?;
451
452 let k4_size = k4.size();
453 if k4_size.len() != 1 || k4_size[0] != y0.size()[0] {
454 anyhow::bail!("Derivative CModule returned tensor of wrong shape at stage k4");
455 }
456
457 let step_div_6 = current_step / 6.0;
459 let y_next = &y + step_div_6 * (&k1 + 2.0 * &k2 + 2.0 * &k3 + &k4);
460
461 validate_finite_tensor(&y_next, "state after RK4 update (NaN/Inf propagating)")?;
463
464 x = &x + current_step;
465 y = y_next;
466
467 all_x.push(x);
468 all_y.push(y.copy());
469
470 step_count += 1;
471
472 let y_norm = y.f_norm()?.f_double_value(&[])?;
473
474 if !warned_large_norm && y_norm > 1e10 {
475 warn!(
476 "RK4: solution norm exceeded {:.1e} at x={:.3e}; the solution may be diverging.",
477 1e10, x
478 );
479 warned_large_norm = true;
480 }
481
482 if !warned_many_steps && step_count >= 100_000 {
483 warn!(
484 "RK4: reached {} steps; consider increasing step size or switching to an adaptive solver",
485 step_count
486 );
487 warned_many_steps = true;
488 }
489 }
490
491 Ok((
492 Tensor::f_from_slice(&all_x)?
493 .to_kind(kind)
494 .to_device(device),
495 Tensor::f_stack(&all_y, 0)?,
496 ))
497}
498
499fn solve_implicit_euler(
501 f: tch::CModule,
502 x_span: (f64, f64),
503 y0: Tensor,
504 step: f64,
505 optimizer: &dyn optimizers::Optimizer,
506) -> anyhow::Result<(Tensor, Tensor)> {
507 let device = y0.device();
508 let kind = y0.kind();
509
510 let x_start = x_span.0;
511 let x_end = x_span.1;
512
513 let mut x = x_start;
514 let mut y = y0.copy();
515
516 let mut all_x = vec![x];
517 let mut all_y = vec![y.copy()];
518
519 let mut current_step = step;
520 let mut step_count: u64 = 0;
521
522 let mut warned_large_norm = false;
523 let mut warned_many_steps = false;
524
525 while x < x_end {
526 let remaining = x_end - x;
527 if remaining < current_step {
528 current_step = remaining;
529 }
530
531 let x_next = &x + current_step;
532 let y_prev = y.copy();
533
534 let f_next_fn = |y_next: &Tensor| {
536 let f_next = f
537 .forward_ts(&[
538 Tensor::from(x_next).to_kind(kind).to_device(device),
539 y_next.copy(),
540 ])
541 .unwrap();
542 let y_pred = &y_prev + current_step * &f_next;
543 (y_next - &y_pred).pow_tensor_scalar(2).sum(y_next.kind())
544 };
545
546 let initial_guess = &y_prev.detach()
548 + current_step
549 * f.forward_ts(&[&Tensor::from(x).to_kind(kind).to_device(device), &y_prev])?;
550
551 let y_next = optimizer
553 .optimize(&f_next_fn, &initial_guess)
554 .map_err(|err| anyhow!(format!("Implicit solver optimizer failed with: {}", err)))?;
555
556 validate_finite_tensor(
558 &y_next,
559 "state after implicit solver optimization (NaN/Inf)",
560 )?;
561
562 y = y_next.copy();
563 x = x_next;
564
565 all_x.push(x);
566 all_y.push(y.copy());
567
568 step_count += 1;
569
570 let y_norm = y.f_norm()?.f_double_value(&[])?;
571
572 if !warned_large_norm && y_norm > 1e10 {
573 warn!(
574 "ImplicitEuler: solution norm exceeded {:.1e} at x={:.3e}; the solution may be diverging.",
575 1e10, x
576 );
577 warned_large_norm = true;
578 }
579
580 if !warned_many_steps && step_count >= 100_000 {
581 warn!(
582 "ImplicitEuler: reached {} steps; consider increasing step size",
583 step_count
584 );
585 warned_many_steps = true;
586 }
587 }
588
589 Ok((
590 Tensor::f_from_slice(&all_x)?
591 .to_kind(kind)
592 .to_device(device),
593 Tensor::f_stack(&all_y, 0)?,
594 ))
595}
596
597fn solve_glrk4(
599 f: tch::CModule,
600 x_span: (f64, f64),
601 y0: Tensor,
602 step: f64,
603 optimizer: &dyn optimizers::Optimizer,
604) -> anyhow::Result<(Tensor, Tensor)> {
605 let device = y0.device();
606 let kind = y0.kind();
607
608 let x_start = x_span.0;
609 let x_end = x_span.1;
610
611 let mut x = x_start;
612 let mut y = y0.copy();
613 let y_length = y.size()[0];
614
615 let mut all_x = vec![x];
616 let mut all_y = vec![y.copy()];
617
618 let mut current_step = step;
619 let mut step_count: u64 = 0;
620
621 let mut warned_large_norm = false;
622 let mut warned_many_steps = false;
623
624 while x < x_end {
625 let remaining = x_end - x;
626 if remaining < current_step {
627 current_step = remaining;
628 }
629
630 let k = f.forward_ts(&[Tensor::from(x).to_kind(kind).to_device(device), y.copy()])?;
631 validate_finite_tensor(&k, "GLRK4 initial derivative")?;
632
633 let k_size = k.size();
634 if k_size.len() != 1 || k_size[0] != y0.size()[0] {
635 anyhow::bail!("Derivative CModule returned tensor of wrong shape in GLRK4");
636 }
637
638 const C1: f64 = 0.2113248654f64;
639 const C2: f64 = 0.7886751346f64;
640 const A11: f64 = 0.25;
641 const A12: f64 = -0.03867513459f64;
642 const A21: f64 = 0.5386751346f64;
643 const A22: f64 = 0.25;
644
645 let first_k1k2_guess = Tensor::f_cat(
647 &[
648 f.forward_ts(&[
649 Tensor::from(x + C1 * current_step)
650 .to_kind(kind)
651 .to_device(device),
652 &y + C1 * current_step * &k,
653 ])?,
654 f.forward_ts(&[
655 Tensor::from(x + C2 * current_step)
656 .to_kind(kind)
657 .to_device(device),
658 &y + C2 * current_step * &k,
659 ])?,
660 ],
661 0,
662 )?;
663
664 let loss_fn = |k1k2_guess: &Tensor| {
666 let diff1 = k1k2_guess.i(0..y_length)
667 - f.forward_ts(&[
668 Tensor::from(x + C1 * current_step)
669 .to_kind(kind)
670 .to_device(device),
671 &y + (A11 * k1k2_guess.i(0..y_length)
672 + A12 * k1k2_guess.i(y_length..2 * y_length))
673 * current_step,
674 ])
675 .unwrap();
676 let diff2 = k1k2_guess.i(y_length..2 * y_length)
677 - f.forward_ts(&[
678 Tensor::from(x + C2 * current_step)
679 .to_kind(kind)
680 .to_device(device),
681 &y + (A21 * k1k2_guess.i(0..y_length)
682 + A22 * k1k2_guess.i(y_length..2 * y_length))
683 * current_step,
684 ])
685 .unwrap();
686 diff1.dot(&diff1) + diff2.dot(&diff2)
687 };
688
689 let k1k2 = optimizer
691 .optimize(&loss_fn, &first_k1k2_guess)
692 .map_err(|err| anyhow!(format!("GLRK4 optimizer failed with: {}", err)))?;
693
694 validate_finite_tensor(
696 &k1k2,
697 "GLRK4 stage coefficients after optimization (NaN/Inf)",
698 )?;
699
700 x = x + current_step;
702 y = &y
703 + current_step
704 * (0.5 * k1k2.f_i(0..y_length)? + 0.5 * k1k2.f_i(y_length..2 * y_length)?);
705
706 validate_finite_tensor(&y, "state after GLRK4 update (NaN/Inf propagating)")?;
708
709 all_x.push(x);
710 all_y.push(y.copy());
711
712 step_count += 1;
713
714 let y_norm = y.f_norm()?.f_double_value(&[])?;
715
716 if !warned_large_norm && y_norm > 1e10 {
717 warn!(
718 "GLRK4: solution norm exceeded {:.1e} at x={:.3e}; the solution may be diverging.",
719 1e10, x
720 );
721 warned_large_norm = true;
722 }
723
724 if !warned_many_steps && step_count >= 100_000 {
725 warn!(
726 "GLRK4: reached {} steps; consider increasing step size",
727 step_count
728 );
729 warned_many_steps = true;
730 }
731 }
732
733 Ok((
734 Tensor::f_from_slice(&all_x)?
735 .to_kind(kind)
736 .to_device(device),
737 Tensor::f_stack(&all_y, 0)?,
738 ))
739}
740
741fn rkf45_step(
743 f: &tch::CModule,
744 x: f64,
745 y: &Tensor,
746 step: f64,
747 device: tch::Device,
748 kind: tch::Kind,
749 y0_length: i64,
750) -> anyhow::Result<(Tensor, Tensor)> {
751 let k1 = f.forward_ts(&[Tensor::from(x).to_kind(kind).to_device(device), y.copy()])?;
753 validate_finite_tensor(&k1, "RKF45 stage k1")?;
754
755 let k1_size = k1.size();
756 if k1_size.len() != 1 || k1_size[0] != y0_length {
757 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
758 }
759
760 let x_step = x + 0.25 * step;
762 let y_step: Tensor = y + 0.25 * &step * &k1;
763 validate_finite_tensor(&y_step, "RKF45 intermediate state for k2")?;
764
765 let k2 = f.forward_ts(&[Tensor::from(x_step).to_kind(kind).to_device(device), y_step])?;
766 validate_finite_tensor(&k2, "RKF45 stage k2")?;
767
768 let k2_size = k2.size();
769 if k2_size.len() != 1 || k2_size[0] != y0_length {
770 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
771 }
772
773 let x_step = x + 0.375 * step;
775 let y_step: Tensor = y + (0.09375 * &step * &k1) + (0.28125 * &step * &k2);
776 validate_finite_tensor(&y_step, "RKF45 intermediate state for k3")?;
777
778 let k3 = f.forward_ts(&[Tensor::from(x_step).to_kind(kind).to_device(device), y_step])?;
779 validate_finite_tensor(&k3, "RKF45 stage k3")?;
780
781 let k3_size = k3.size();
782 if k3_size.len() != 1 || k3_size[0] != y0_length {
783 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
784 }
785
786 let x_step = x + (12.0 / 13.0) * step;
788 let y_step: Tensor = y
789 + (1932.0 / 2197.0 * &step * &k1)
790 + (-7200.0 / 2197.0 * &step * &k2)
791 + (7296.0 / 2197.0 * &step * &k3);
792 validate_finite_tensor(&y_step, "RKF45 intermediate state for k4")?;
793
794 let k4 = f.forward_ts(&[Tensor::from(x_step).to_kind(kind).to_device(device), y_step])?;
795 validate_finite_tensor(&k4, "RKF45 stage k4")?;
796
797 let k4_size = k4.size();
798 if k4_size.len() != 1 || k4_size[0] != y0_length {
799 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
800 }
801
802 let x_step = x + step;
804 let y_step: Tensor = y
805 + (439.0 / 216.0 * &step * &k1)
806 + (-8.0 * &step * &k2)
807 + (3680.0 / 513.0 * &step * &k3)
808 + (-845.0 / 4104.0 * &step * &k4);
809 validate_finite_tensor(&y_step, "RKF45 intermediate state for k5")?;
810
811 let k5 = f.forward_ts(&[Tensor::from(x_step).to_kind(kind).to_device(device), y_step])?;
812 validate_finite_tensor(&k5, "RKF45 stage k5")?;
813
814 let k5_size = k5.size();
815 if k5_size.len() != 1 || k5_size[0] != y0_length {
816 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
817 }
818
819 let x_step = x + 0.5 * step;
821 let y_step: Tensor = y
822 + (-8.0 / 27.0 * &step * &k1)
823 + (2.0 * &step * &k2)
824 + (-3544.0 / 2565.0 * &step * &k3)
825 + (1859.0 / 4104.0 * &step * &k4)
826 + (-11.0 / 40.0 * &step * &k5);
827 validate_finite_tensor(&y_step, "RKF45 intermediate state for k6")?;
828
829 let k6 = f.forward_ts(&[Tensor::from(x_step).to_kind(kind).to_device(device), y_step])?;
830 validate_finite_tensor(&k6, "RKF45 stage k6")?;
831
832 let k6_size = k6.size();
833 if k6_size.len() != 1 || k6_size[0] != y0_length {
834 anyhow::bail!("Derivative CModule returned tensor of bad shape in RKF45");
835 }
836
837 let next_y4: Tensor = y + step
839 * ((25.0 / 216.0 * &k1)
840 + (1408.0 / 2565.0 * &k3)
841 + (2197.0 / 4104.0 * &k4)
842 + (-1.0 / 5.0 * &k5));
843
844 let next_y5: Tensor = y + step
845 * ((16.0 / 135.0 * &k1)
846 + (6656.0 / 12825.0 * &k3)
847 + (28561.0 / 56430.0 * &k4)
848 + (-9.0 / 50.0 * &k5)
849 + (2.0 / 55.0 * &k6));
850
851 Ok((next_y4, next_y5))
852}
853
854fn solve_rkf45(
856 f: tch::CModule,
857 x_span: (f64, f64),
858 y0: Tensor,
859 rtol: f64,
860 atol: f64,
861 min_step: f64,
862 safety_factor: f64,
863) -> anyhow::Result<(Tensor, Tensor)> {
864 let device = y0.device();
865 let kind = y0.kind();
866
867 let x_start = x_span.0;
868 let x_end = x_span.1;
869
870 let mut x = x_start;
871 let mut y = y0.copy();
872
873 let mut all_x = vec![x];
874 let mut all_y = vec![y.copy()];
875
876 let mut step = (x_end - x_start) * 0.1;
877
878 let mut consecutive_rejections: u32 = 0;
879 let mut total_rejections: u32 = 0;
880 let mut total_accepted: u32 = 0;
881
882 let mut warned_consec_rej = false;
883 let mut warned_total_rej = false;
884 let mut warned_tiny_step = false;
885
886 const MAX_GROWTH: f64 = 5.;
887
888 while x < x_end {
889 let (next_y4, next_y5) = rkf45_step(&f, x, &y, step, device, kind, y0.size()[0])?;
890
891 let d = (&next_y4 - &next_y5).f_abs()?;
893 validate_finite_tensor(&d, "RKF45 error estimate difference")?;
894
895 let e = next_y5.f_abs()? * rtol + atol;
896 validate_finite_tensor(&e, "RKF45 error tolerance combination")?;
897
898 let d_min = d.f_min()?.f_double_value(&[])?;
900 let d_max = d.f_max()?.f_double_value(&[])?;
901 let e_min = e.f_min()?.f_double_value(&[])?;
902 let e_max = e.f_max()?.f_double_value(&[])?;
903 println!(
904 "x={:.17e}, step={:.17e}, d=[{:.3e},{:.3e}], e=[{:.3e},{:.3e}]",
905 x, step, d_min, d_max, e_min, e_max
906 );
907
908 let alpha = (e / d)
910 .f_pow_tensor_scalar(0.2)?
911 .f_min()?
912 .f_double_value(&[])?;
913
914 let condition = (safety_factor * alpha).clamp(0f64, MAX_GROWTH);
915
916 if condition < 1f64 {
917 consecutive_rejections += 1;
919 total_rejections += 1;
920
921 if !warned_consec_rej && consecutive_rejections >= 20 {
923 warn!(
924 "RKF45: {} consecutive rejected steps at x={:.3e}, step={:.3e}; problem may be stiff",
925 consecutive_rejections, x, step
926 );
927 warned_consec_rej = true;
928 }
929
930 if !warned_total_rej && total_rejections >= 1000 {
932 warn!(
933 "RKF45: {} total rejected steps ({} accepted) at x={:.3e}; integration is inefficient",
934 total_rejections, total_accepted, x
935 );
936 warned_total_rej = true;
937 }
938
939 if !warned_tiny_step && step < min_step * 10.0 {
941 warn!(
942 "RKF45: required very small step {:.3e} at x={:.3e} (min_step={:.3e}); solution may be inaccurate or problem is stiff",
943 step, x, min_step
944 );
945 warned_tiny_step = true;
946 }
947
948 step = step * condition;
949 validate_finite_scalar(step, "RKF45 reduced step size")?;
950
951 if step < min_step {
953 return Err(anyhow!("Required step is smaller than minimal step"));
954 }
955 } else {
956 consecutive_rejections = 0;
958 total_accepted += 1;
959
960 let remaining = x_end - x;
962 if remaining < step {
963 step = remaining;
964 let (_next_y4, next_y5) = rkf45_step(&f, x, &y, step, device, kind, y0.size()[0])?;
965 y = next_y5;
966 x = x_end;
967 all_x.push(x);
968 all_y.push(y.copy());
969 break;
970 }
971
972 y = next_y5;
973 x = &x + &step;
974
975 validate_finite_tensor(&y, "RKF45 accepted state (NaN/Inf)")?;
977 validate_finite_scalar(x, "RKF45 updated integration variable")?;
978
979 all_x.push(x);
980 all_y.push(y.copy());
981
982 step = step * condition;
983 validate_finite_scalar(step, "RKF45 next step size")?;
984 }
985 }
986
987 if total_rejections > total_accepted * 2 && total_accepted > 0 {
989 warn!(
990 "RKF45: integration completed with {} rejected and {} accepted steps; consider relaxing tolerances or using an implicit solver for stiff problems",
991 total_rejections, total_accepted
992 );
993 }
994
995 Ok((
996 Tensor::f_from_slice(&all_x)?
997 .to_kind(kind)
998 .to_device(device),
999 Tensor::f_stack(&all_y, 0)?,
1000 ))
1001}
1002
1003fn solve_row1(
1005 f: tch::CModule,
1006 x_span: (f64, f64),
1007 y0: Tensor,
1008 step: f64,
1009) -> anyhow::Result<(Tensor, Tensor)> {
1010 let device = y0.device();
1011 let kind = y0.kind();
1012
1013 let x_start = x_span.0;
1014 let x_end = x_span.1;
1015
1016 let mut x = x_start;
1017 let mut y = y0.copy();
1018
1019 let mut all_x = vec![x];
1020 let mut all_y = vec![y.copy()];
1021
1022 let mut step_count: u64 = 0;
1023
1024 let mut warned_large_matrix = false;
1025 let mut warned_large_inverse = false;
1026 let mut warned_large_norm = false;
1027 let mut warned_many_steps = false;
1028
1029 while x < x_end {
1030 let remaining = x_end - x;
1031 let mut current_step = step;
1032 if remaining < step {
1033 current_step = remaining;
1034 }
1035
1036 let x_prev = x;
1037 let y_prev = y.copy();
1038
1039 let jacobian = compute_jacobian(
1041 |y| {
1042 f.forward_ts(&[
1043 Tensor::from(x_prev).to_kind(kind).to_device(device),
1044 y.copy(),
1045 ])
1046 .unwrap()
1047 },
1048 &y_prev,
1049 )?;
1050
1051 validate_finite_tensor(&jacobian, "Jacobian matrix in ROW1 (NaN/Inf)")?;
1053
1054 let f_current = f.forward_ts(&[
1056 Tensor::from(x_prev).to_kind(kind).to_device(device),
1057 y_prev.copy(),
1058 ])?;
1059
1060 validate_finite_tensor(&f_current, "derivative function output in ROW1")?;
1061
1062 let f_current_size = f_current.size();
1063 let f_current_rank = f_current_size.len();
1064 if f_current_rank != 1 {
1065 anyhow::bail!(
1066 "Derivative CModule returned tensor of bad rank {}.",
1067 f_current_rank
1068 );
1069 }
1070 if f_current_size[0] != y0.size()[0] {
1071 anyhow::bail!(
1072 "Derivative CModule returned vector of bad length {}.",
1073 f_current_size[0]
1074 );
1075 }
1076
1077 let n = jacobian.size()[0];
1079 let eye = Tensor::f_eye(n, (jacobian.kind(), jacobian.device()))?;
1080 let step_j = current_step * &jacobian;
1081 let matrix_to_invert = eye - step_j;
1082
1083 let matrix_norm = matrix_to_invert.f_norm()?.f_double_value(&[])?;
1085 if !warned_large_matrix && matrix_norm > 1e12 {
1086 warn!(
1087 "ROW1: linear system matrix has large norm {:.3e} at x={:.3e}; solution may be unstable",
1088 matrix_norm, x_prev
1089 );
1090 warned_large_matrix = true;
1091 }
1092
1093 let inv_matrix = matrix_to_invert.f_inverse()?;
1094
1095 validate_finite_tensor(&inv_matrix, "inverse matrix (I - h*J)^(-1) in ROW1")?;
1096
1097 let inv_norm = inv_matrix.f_norm()?.f_double_value(&[])?;
1099 if !warned_large_inverse && inv_norm > 1e10 {
1100 warn!(
1101 "ROW1: inverse matrix has large norm {:.3e} at x={:.3e}; Jacobian may be ill-conditioned",
1102 inv_norm, x_prev
1103 );
1104 warned_large_inverse = true;
1105 }
1106
1107 let delta_y = inv_matrix.f_matmul(&f_current)?;
1108 validate_finite_tensor(&delta_y, "Newton correction step in ROW1")?;
1109
1110 let y_next = y_prev + current_step * delta_y;
1111
1112 validate_finite_tensor(&y_next, "state after ROW1 update (NaN/Inf)")?;
1114
1115 x = &x_prev + current_step;
1116 validate_finite_scalar(x, "ROW1 updated integration variable")?;
1117
1118 y = y_next.detach().copy();
1119
1120 all_x.push(x);
1121 all_y.push(y.copy());
1122
1123 step_count += 1;
1124
1125 let y_norm = y.f_norm()?.f_double_value(&[])?;
1126
1127 if !warned_large_norm && y_norm > 1e10 {
1128 warn!(
1129 "ROW1: solution norm exceeded {:.1e} at x={:.3e}; the solution may be diverging.",
1130 1e10, x
1131 );
1132 warned_large_norm = true;
1133 }
1134
1135 if !warned_many_steps && step_count >= 100_000 {
1136 warn!(
1137 "ROW1: reached {} steps; consider increasing step size",
1138 step_count
1139 );
1140 warned_many_steps = true;
1141 }
1142 }
1143
1144 Ok((
1145 Tensor::f_from_slice(&all_x)?
1146 .to_kind(kind)
1147 .to_device(device),
1148 Tensor::f_stack(&all_y, 0)?,
1149 ))
1150}
1151
1152fn compute_jacobian<F>(f: F, x: &Tensor) -> anyhow::Result<Tensor>
1154where
1155 F: Fn(&Tensor) -> Tensor,
1156{
1157 if x.dim() != 1 {
1158 return Err(anyhow!(
1159 "Jacobian input tensor must be one-dimensional, got {} dimensions",
1160 x.dim()
1161 ));
1162 }
1163
1164 let x_with_grad = x.detach().copy().set_requires_grad(true);
1165
1166 let y = f(&x_with_grad);
1167
1168 if y.dim() != 1 {
1169 return Err(anyhow!(
1170 "Jacobian output tensor must be one-dimensional, got {} dimensions",
1171 y.dim()
1172 ));
1173 }
1174
1175 if y.isfinite().f_all()?.f_int64_value(&[])? == 0 {
1176 return Err(anyhow!(
1177 "Jacobian function returned tensor containing non-finite values"
1178 ));
1179 }
1180
1181 let y_size = y.size()[0];
1182 let mut grads = Vec::with_capacity(y_size as usize);
1183
1184 for i in 0..y_size {
1185 let yi = y.i(i);
1186
1187 let grad = Tensor::f_run_backward(&[yi], &[&x_with_grad], true, false)?
1188 .first()
1189 .ok_or_else(|| anyhow!("Failed to compute Jacobian gradient"))?
1190 .copy();
1191
1192 if grad.size() != x.size() {
1193 return Err(anyhow!(
1194 "Jacobian gradient has shape {:?}, expected shape {:?}",
1195 grad.size(),
1196 x.size()
1197 ));
1198 }
1199
1200 if grad.isfinite().f_all()?.f_int64_value(&[])? == 0 {
1201 return Err(anyhow!("Jacobian computation produced non-finite values"));
1202 }
1203
1204 grads.push(grad);
1205 }
1206
1207 let jacobian = Tensor::f_stack(&grads, 0)?;
1208
1209 if jacobian.size() != vec![y_size, x.size()[0]] {
1210 return Err(anyhow!(
1211 "Jacobian has shape {:?}, expected shape [{}, {}]",
1212 jacobian.size(),
1213 y_size,
1214 x.size()[0]
1215 ));
1216 }
1217
1218 Ok(jacobian)
1219}