1use anyhow::anyhow;
2use std::fmt;
3use std::sync::Arc;
4use tch::IndexOp;
5use tch::Tensor;
6
7pub mod optimizers;
8
9pub enum Solver {
10 Euler {
11 step: f64,
12 },
13 RK4 {
14 step: f64,
15 },
16 ImplicitEuler {
17 step: f64,
18 optimizer: Arc<dyn optimizers::Optimizer>,
19 },
20 GLRK4 {
21 step: f64,
22 optimizer: Arc<dyn optimizers::Optimizer>,
23 },
24 RKF45 {
25 rtol: f64,
26 atol: f64,
27 min_step: f64,
28 safety_factor: f64,
29 },
30 ROW1 {
31 step: f64,
32 },
33}
34
35impl Solver {
36 pub fn solve(
37 &self,
38 f: tch::CModule,
39 x_span: (f64, f64),
40 y0: Tensor,
41 ) -> anyhow::Result<(Tensor, Tensor)> {
42 if y0.size().len() != 1 {
43 return Err(anyhow!(
44 "y0 must be a one-dimensional tensor but it has {} dimensions",
45 y0.size().len()
46 ));
47 }
48 if y0.kind() != tch::Kind::Double
49 && y0.kind() != tch::Kind::Float
50 && y0.kind() != tch::Kind::BFloat16
51 && y0.kind() != tch::Kind::Half
52 {
53 return Err(anyhow!("y0 is of unsupported kind {:?}", y0.kind()));
54 }
55
56 match self {
57 Self::Euler { step } => solve_euler(f, x_span, y0, *step),
58 Self::RK4 { step } => solve_rk4(f, x_span, y0, *step),
59 Self::ImplicitEuler { step, optimizer } => {
60 solve_implicit_euler(f, x_span, y0, *step, optimizer.as_ref())
61 }
62 Self::GLRK4 { step, optimizer } => {
63 solve_glrk4(f, x_span, y0, *step, optimizer.as_ref())
64 }
65 Self::RKF45 {
66 rtol,
67 atol,
68 min_step,
69 safety_factor,
70 } => solve_rkf45(f, x_span, y0, *rtol, *atol, *min_step, *safety_factor),
71 Self::ROW1 { step } => solve_row1(f, x_span, y0, *step),
72 }
73 }
74}
75
76impl fmt::Display for Solver {
77 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
78 match self {
79 Solver::Euler { step } => write!(f, "Euler(step={})", step),
80 Solver::RK4 { step } => write!(f, "RK4(step={})", step),
81 Solver::ImplicitEuler { step, optimizer } => {
82 write!(f, "ImplicitEuler(step={}, optimizer={})", step, optimizer)
83 }
84 Solver::GLRK4 { step, optimizer } => {
85 write!(f, "GLRK4(step={}, optimizer={})", step, optimizer)
86 }
87 Solver::RKF45 {
88 rtol,
89 atol,
90 min_step,
91 safety_factor,
92 } => write!(
93 f,
94 "RKF45(rtol={}, atol={}, min_step={}, safety_factor={})",
95 rtol, atol, min_step, safety_factor
96 ),
97 Solver::ROW1 { step } => write!(f, "ROW1(step={})", step),
98 }
99 }
100}
101
102fn solve_euler(
104 f: tch::CModule,
105 x_span: (f64, f64),
106 y0: Tensor,
107 step: f64,
108) -> anyhow::Result<(Tensor, Tensor)> {
109 let device = y0.device();
110 let kind = y0.kind();
111
112 let x_start = x_span.0;
113 let x_end = x_span.1;
114
115 let mut x = x_start;
116 let mut y = y0.unsqueeze(0);
117
118 let mut all_x = vec![x];
119 let mut all_y = vec![y.copy()];
120
121 let mut current_step = step;
122 while x < x_end {
123 let remaining = x_end - x;
124 if remaining < current_step {
125 current_step = remaining;
126 }
127
128 let dy = f.forward_ts(&[
129 Tensor::from(x).to_kind(kind).to_device(device),
130 y.squeeze().copy(),
131 ])?;
132 let dy_rank = dy.size().len();
133 if dy_rank != 1 {
134 anyhow::bail!(
135 "Derivative CModule returned tensor of bad rank {}.",
136 dy_rank
137 );
138 }
139
140 y = &y + current_step * &dy;
141 x = &x + current_step;
142
143 all_x.push(x);
144 all_y.push(y.copy());
145 }
146
147 Ok((
148 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
149 Tensor::cat(&all_y, 0),
150 ))
151}
152
153fn solve_rk4(
155 f: tch::CModule,
156 x_span: (f64, f64),
157 y0: Tensor,
158 step: f64,
159) -> anyhow::Result<(Tensor, Tensor)> {
160 let device = y0.device();
161 let kind = y0.kind();
162
163 let x_start = x_span.0;
164 let x_end = x_span.1;
165
166 let mut x = x_start;
167 let mut y = y0.unsqueeze(0);
168
169 let mut all_x = vec![x];
170 let mut all_y = vec![y.copy()];
171
172 let mut current_step = step;
173 while x < x_end {
174 let remaining = x_end - x;
175 if remaining < current_step {
176 current_step = remaining;
177 }
178
179 let k1 = f.forward_ts(&[
180 Tensor::from(x).to_kind(kind).to_device(device),
181 y.squeeze().copy(),
182 ])?;
183 let k1_rank = k1.size().len();
184 if k1_rank != 1 {
185 anyhow::bail!(
186 "Derivative CModule returned tensor of bad rank {}.",
187 k1_rank
188 );
189 }
190
191 let x_half = x + 0.5 * current_step;
192 let y_half: Tensor = &y + 0.5 * current_step * &k1;
193 let k2 = f.forward_ts(&[
194 Tensor::from(x_half).to_kind(kind).to_device(device),
195 y_half.squeeze(),
196 ])?;
197 let k2_rank = k2.size().len();
198 if k2_rank != 1 {
199 anyhow::bail!(
200 "Derivative CModule returned tensor of bad rank {}.",
201 k2_rank
202 );
203 }
204
205 let x_half_again = x + 0.5 * current_step;
206 let y_half_again: Tensor = &y + 0.5 * current_step * &k2;
207 let k3 = f.forward_ts(&[
208 Tensor::from(x_half_again).to_kind(kind).to_device(device),
209 y_half_again.squeeze(),
210 ])?;
211 let k3_rank = k3.size().len();
212 if k3_rank != 1 {
213 anyhow::bail!(
214 "Derivative CModule returned tensor of bad rank {}.",
215 k3_rank
216 );
217 }
218
219 let x_full = x + current_step;
220 let y_full = &y + current_step * &k3;
221 let k4 = f.forward_ts(&[
222 Tensor::from(x_full).to_kind(kind).to_device(device),
223 y_full.squeeze(),
224 ])?;
225 let k4_rank = k4.size().len();
226 if k4_rank != 1 {
227 anyhow::bail!(
228 "Derivative CModule returned tensor of bad rank {}.",
229 k4_rank
230 );
231 }
232
233 let step_div_6 = current_step / 6.0;
234 let y_next = &y + step_div_6 * (&k1 + 2.0 * &k2 + 2.0 * &k3 + &k4);
235
236 x = &x + current_step;
237 y = y_next;
238
239 all_x.push(x);
240 all_y.push(y.copy());
241 }
242
243 Ok((
244 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
245 Tensor::cat(&all_y, 0),
246 ))
247}
248
249fn solve_implicit_euler(
251 f: tch::CModule,
252 x_span: (f64, f64),
253 y0: Tensor,
254 step: f64,
255 optimizer: &dyn optimizers::Optimizer,
256) -> anyhow::Result<(Tensor, Tensor)> {
257 let device = y0.device();
258 let kind = y0.kind();
259
260 let x_start = x_span.0;
261 let x_end = x_span.1;
262
263 let mut x = x_start;
264 let mut y = y0.unsqueeze(0);
265
266 let mut all_x = vec![x];
267 let mut all_y = vec![y.copy()];
268
269 let mut current_step = step;
270 while x < x_end {
271 let remaining = x_end - x;
272 if remaining < current_step {
273 current_step = remaining;
274 }
275
276 let x_next = &x + current_step;
277 let y_prev = y.copy();
278
279 let y_next = optimizer
280 .optimize(
281 &|y_next: &Tensor| {
282 let f_next = f
283 .forward_ts(&[
284 Tensor::from(x_next).to_kind(kind).to_device(device),
285 y_next.squeeze().copy(),
286 ])
287 .unwrap();
288 let y_pred = &y_prev.squeeze() + current_step * &f_next;
289 (y_next - &y_pred).pow_tensor_scalar(2).sum(y_next.kind())
290 },
291 &(&y_prev.detach().squeeze()
292 + current_step
293 * f.forward_ts(&[
294 &Tensor::from(x).to_kind(kind).to_device(device),
295 &y_prev.squeeze(),
296 ])?),
297 )
298 .map_err(|err| anyhow!(format!("Optimizer failed with: {}", err)))?;
299
300 y = y_next.unsqueeze(0);
301 x = x_next;
302
303 all_x.push(x);
304 all_y.push(y.copy());
305 }
306
307 Ok((
308 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
309 Tensor::cat(&all_y, 0),
310 ))
311}
312
313fn solve_glrk4(
315 f: tch::CModule,
316 x_span: (f64, f64),
317 y0: Tensor,
318 step: f64,
319 optimizer: &dyn optimizers::Optimizer,
320) -> anyhow::Result<(Tensor, Tensor)> {
321 let device = y0.device();
322 let kind = y0.kind();
323
324 let x_start = x_span.0;
325 let x_end = x_span.1;
326
327 let mut x = x_start;
328 let mut y = y0.unsqueeze(0);
329
330 let mut all_x = vec![x];
331 let mut all_y = vec![y.copy()];
332
333 let mut current_step = step;
334 while x < x_end {
335 let remaining = x_end - x;
336 if remaining < current_step {
337 current_step = remaining;
338 }
339
340 let k = f.forward_ts(&[
341 Tensor::from(x).to_kind(kind).to_device(device),
342 y.squeeze().copy(),
343 ])?;
344 let k_rank = k.size().len();
345 if k_rank != 1 {
346 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k_rank);
347 }
348
349 const C1: f64 = 0.2113248654f64;
350 const C2: f64 = 0.7886751346f64;
351 const A11: f64 = 0.25;
352 const A12: f64 = -0.03867513459f64;
353 const A21: f64 = 0.5386751346f64;
354 const A22: f64 = 0.25;
355
356 let first_k1k2_guess = Tensor::cat(
357 &[
358 f.forward_ts(&[
359 Tensor::from(x + C1 * current_step)
360 .to_kind(kind)
361 .to_device(device),
362 &y.squeeze() + C1 * current_step * &k,
363 ])?,
364 f.forward_ts(&[
365 Tensor::from(x + C2 * current_step)
366 .to_kind(kind)
367 .to_device(device),
368 &y.squeeze() + C2 * current_step * &k,
369 ])?,
370 ],
371 0,
372 );
373 let k1k2 = optimizer
374 .optimize(
375 &|k1k2_guess| {
376 let diff1 = k1k2_guess.i(0..=1)
377 - f.forward_ts(&[
378 Tensor::from(x + C1 * current_step)
379 .to_kind(kind)
380 .to_device(device),
381 &y.squeeze()
382 + (A11 * k1k2_guess.i(0..=1) + A12 * k1k2_guess.i(2..=3))
383 * current_step,
384 ])
385 .unwrap();
386 let diff2 = k1k2_guess.i(2..=3)
387 - f.forward_ts(&[
388 Tensor::from(x + C2 * current_step)
389 .to_kind(kind)
390 .to_device(device),
391 &y.squeeze()
392 + (A21 * k1k2_guess.i(0..=1) + A22 * k1k2_guess.i(2..=3))
393 * current_step,
394 ])
395 .unwrap();
396
397 diff1.dot(&diff1) + diff2.dot(&diff2)
398 },
399 &first_k1k2_guess,
400 )
401 .map_err(|err| anyhow!(format!("Optimizer failed with: {}", err)))?;
402 assert!(k1k2.size().len() == 1);
403 assert!(k1k2.size()[0] == 4);
404
405 x = x + current_step;
406 y = &y + current_step * (0.5 * k1k2.i(0..=1) + 0.5 * k1k2.i(2..=3));
407
408 all_x.push(x);
409 all_y.push(y.copy());
410 }
411
412 Ok((
413 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
414 Tensor::cat(&all_y, 0),
415 ))
416}
417
418fn solve_rkf45(
420 f: tch::CModule,
421 x_span: (f64, f64),
422 y0: Tensor,
423 rtol: f64,
424 atol: f64,
425 min_step: f64,
426 safety_factor: f64,
427) -> anyhow::Result<(Tensor, Tensor)> {
428 let device = y0.device();
429 let kind = y0.kind();
430
431 let x_start = x_span.0;
432 let x_end = x_span.1;
433
434 let mut x = x_start;
435 let mut y = y0.unsqueeze(0);
436
437 let mut all_x = vec![x];
438 let mut all_y = vec![y.copy()];
439
440 let mut step = (x_end - x_start) * 0.1;
441
442 while x < x_end {
443 let remaining = x_end - x;
444 if remaining < step {
445 step = remaining;
446 }
447
448 let k1 = f.forward_ts(&[
449 Tensor::from(x).to_kind(kind).to_device(device),
450 y.squeeze().copy(),
451 ])?;
452 let k1_rank = k1.size().len();
453 if k1_rank != 1 {
454 anyhow::bail!(
455 "Derivative CModule returned tensor of bad rank {}.",
456 k1_rank
457 );
458 }
459
460 let k2 = {
461 let x_step = x + 0.25 * step;
462 let y_step: Tensor = &y + 0.25 * &step * &k1;
463 let k2_unchecked = f.forward_ts(&[
464 Tensor::from(x_step).to_kind(kind).to_device(device),
465 y_step.squeeze(),
466 ])?;
467 let k2_rank = k2_unchecked.size().len();
468 if k2_rank != 1 {
469 anyhow::bail!(
470 "Derivative CModule returned tensor of bad rank {}.",
471 k2_rank
472 );
473 }
474
475 k2_unchecked
476 };
477
478 let k3 = {
479 let x_step = x + 0.375 * step;
480 let y_step: Tensor = &y + (0.09375 * &step * &k1) + (0.28125 * &step * &k2);
481 let k3_unchecked = f.forward_ts(&[
482 Tensor::from(x_step).to_kind(kind).to_device(device),
483 y_step.squeeze(),
484 ])?;
485 let k3_rank = k3_unchecked.size().len();
486 if k3_rank != 1 {
487 anyhow::bail!(
488 "Derivative CModule returned tensor of bad rank {}.",
489 k3_rank
490 );
491 }
492
493 k3_unchecked
494 };
495
496 let k4 = {
497 let x_step = x + (12.0 / 13.0) * step;
498 let y_step: Tensor = &y
499 + (1932.0 / 2197.0 * &step * &k1)
500 + (-7200.0 / 2197.0 * &step * &k2)
501 + (7296.0 / 2197.0 * &step * &k3);
502 let k4_unchecked = f.forward_ts(&[
503 Tensor::from(x_step).to_kind(kind).to_device(device),
504 y_step.squeeze(),
505 ])?;
506 let k4_rank = k4_unchecked.size().len();
507 if k4_rank != 1 {
508 anyhow::bail!(
509 "Derivative CModule returned tensor of bad rank {}.",
510 k4_rank
511 );
512 }
513
514 k4_unchecked
515 };
516
517 let k5 = {
518 let x_step = x + step;
519 let y_step: Tensor = &y
520 + (439.0 / 216.0 * &step * &k1)
521 + (-8.0 * &step * &k2)
522 + (3680.0 / 513.0 * &step * &k3)
523 + (-845.0 / 4104.0 * &step * &k4);
524 let k5_unchecked = f.forward_ts(&[
525 Tensor::from(x_step).to_kind(kind).to_device(device),
526 y_step.squeeze(),
527 ])?;
528 let k5_rank = k5_unchecked.size().len();
529 if k5_rank != 1 {
530 anyhow::bail!(
531 "Derivative CModule returned tensor of bad rank {}.",
532 k5_rank
533 );
534 }
535
536 k5_unchecked
537 };
538
539 let k6 = {
540 let x_step = x + 0.5 * step;
541 let y_step: Tensor = &y
542 + (-8.0 / 27.0 * &step * &k1)
543 + (2.0 * &step * &k2)
544 + (-3544.0 / 2565.0 * &step * &k3)
545 + (1859.0 / 4104.0 * &step * &k4)
546 + (-11.0 / 40.0 * &step * &k5);
547 let k6_unchecked = f.forward_ts(&[
548 Tensor::from(x_step).to_kind(kind).to_device(device),
549 y_step.squeeze(),
550 ])?;
551 let k6_rank = k6_unchecked.size().len();
552 if k6_rank != 1 {
553 anyhow::bail!(
554 "Derivative CModule returned tensor of bad rank {}.",
555 k6_rank
556 );
557 }
558
559 k6_unchecked
560 };
561
562 let next_y4: Tensor = &y
563 + step
564 * ((25.0 / 216.0 * &k1)
565 + (1408.0 / 2565.0 * &k3)
566 + (2197.0 / 4104.0 * &k4)
567 + (-1.0 / 5.0 * &k5));
568 let next_y5: Tensor = &y
569 + step
570 * ((16.0 / 135.0 * &k1)
571 + (6656.0 / 12825.0 * &k3)
572 + (28561.0 / 56430.0 * &k4)
573 + (-9.0 / 50.0 * &k5)
574 + (2.0 / 55.0 * &k6));
575
576 let d = (&next_y4 - &next_y5).abs();
577 let e = next_y5.abs() * rtol + atol;
578
579 let alpha = (e / d).sqrt().min().double_value(&[]);
580 let condition = safety_factor * alpha;
581
582 if condition < 1f64 {
583 step = step * condition;
584 if step < min_step {
585 return Err(anyhow!("Required step is smaller than minimal step"));
586 }
587 } else {
588 y = next_y4;
589 x = &x + &step;
590 all_x.push(x);
591 all_y.push(y.copy());
592
593 let new_step = step * condition;
594 let max_step = step * 5.0;
595 step = match new_step.partial_cmp(&max_step) {
596 Some(std::cmp::Ordering::Less) => new_step,
597 Some(_) => max_step,
598 None => return Err(anyhow!("Runtime error")),
599 };
600 }
601 }
602
603 Ok((
604 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
605 Tensor::cat(&all_y, 0),
606 ))
607}
608
609fn solve_row1(
611 f: tch::CModule,
612 x_span: (f64, f64),
613 y0: Tensor,
614 step: f64,
615) -> anyhow::Result<(Tensor, Tensor)> {
616 let device = y0.device();
617 let kind = y0.kind();
618
619 let x_start = x_span.0;
620 let x_end = x_span.1;
621
622 let mut x = x_start;
623 let mut y = y0.unsqueeze(0);
624
625 let mut all_x = vec![x];
626 let mut all_y = vec![y.copy()];
627
628 while x < x_end {
629 let remaining = x_end - x;
630 let mut current_step = step;
631 if remaining < step {
632 current_step = remaining;
633 }
634
635 let x_prev = x;
636 let y_prev = y.copy().squeeze();
637
638 let jacobian = compute_jacobian(
639 |y| {
640 f.forward_ts(&[
641 Tensor::from(x_prev).to_kind(kind).to_device(device),
642 y.copy(),
643 ])
644 .unwrap()
645 .squeeze()
646 },
647 &y_prev,
648 );
649 let f_current = f.forward_ts(&[
650 Tensor::from(x_prev).to_kind(kind).to_device(device),
651 y_prev.copy(),
652 ])?;
653 let f_current_rank = f_current.size().len();
654 if f_current_rank != 1 {
655 anyhow::bail!(
656 "Derivative CModule returned tensor of bad rank {}.",
657 f_current_rank
658 );
659 }
660
661 let n = jacobian.size()[0];
662 let eye = Tensor::eye(n, (tch::Kind::Float, jacobian.device()));
663 let step_j = current_step * &jacobian;
664 let inv_matrix = (eye - step_j).inverse();
665
666 let delta_y = inv_matrix.matmul(&f_current);
667 let y_next = y_prev + current_step * delta_y;
668
669 x = &x_prev + current_step;
670 y = y_next.unsqueeze(0);
671
672 all_x.push(x);
673 all_y.push(y.copy());
674 }
675
676 Ok((
677 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
678 Tensor::cat(&all_y, 0),
679 ))
680}
681
682fn compute_jacobian<F>(f: F, x: &Tensor) -> Tensor
684where
685 F: Fn(&Tensor) -> Tensor,
686{
687 assert_eq!(x.dim(), 1, "x must be 1-dimensional");
688 let mut x_with_grad = x.detach().copy().set_requires_grad(true);
689 let y = f(&x_with_grad);
690 assert_eq!(y.dim(), 1, "y must be 1-dimensional");
691
692 let y_size = y.size()[0];
693 let mut grads = Vec::new();
694
695 for i in 0..y_size {
696 let yi = y.i(i);
697 let grad = Tensor::run_backward(&[yi], &[&x_with_grad], true, false)[0].copy();
698 grads.push(grad.unsqueeze(0));
699 x_with_grad.zero_grad();
700 }
701
702 Tensor::cat(&grads, 0)
703}