mini_ode/
lib.rs

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
102/// Solves ODE using Euler method
103fn 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
153/// Solves ODE using Runge-Kutta 4th order method
154fn 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
249/// Solves ODE using Implicit Euler method with gradient descent optimization
250fn 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
313/// Solves ODE using Gauss-Legendre-Runge-Kutta 4th order method
314fn 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
418/// Solves ODE using Runge-Kutta-Fehlberg 45 adaptive method
419fn 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
609/// Solves ODE using first-order Rosenbrock method (Row1)
610fn 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
682/// Computes the Jacobian matrix of a function f at point x
683fn 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}