mini_ode/
lib.rs

1use anyhow::anyhow;
2use tch::IndexOp;
3use tch::Tensor;
4use std::sync::Arc;
5
6pub mod optimizers;
7
8pub enum Solver {
9    Euler { step: f64 },
10    RK4 { step: f64 },
11    ImplicitEuler { step: f64, optimizer: Arc<dyn optimizers::Optimizer> },
12    GLRK4 { step: f64, optimizer: Arc<dyn optimizers::Optimizer> },
13    RKF45 { rtol: f64, atol: f64, min_step: f64, safety_factor: f64 },
14    ROW1 { step: f64 }
15}
16
17impl Solver {
18    pub fn solve(
19        &self,
20        f: tch::CModule,
21        x_span: Tensor,
22        y0: Tensor
23    ) -> anyhow::Result<(Tensor, Tensor)> {
24        if x_span.size() != [2] {
25            return Err(anyhow!("x_span must be of shape [2] but it has shape {:?}", x_span.size().as_slice()));
26        }
27        if y0.size().len() != 1 {
28            return Err(anyhow!("y0 must be a one-dimensional tensor but it has {} dimensions", y0.size().len()));
29        }
30        if x_span.device() != y0.device() {
31            return Err(anyhow!("x_span and y0 must reside on the same device. Device of x_span is {:?}. Device of y0 is {:?}", x_span.device(), y0.device()));
32        }
33        if x_span.kind() != tch::Kind::Double && x_span.kind() != tch::Kind::Float && x_span.kind() != tch::Kind::BFloat16 && x_span.kind() != tch::Kind::Half {
34            return Err(anyhow!("x_span is of unsupported kind {:?}", x_span.kind()));
35        }
36        if y0.kind() != tch::Kind::Double && y0.kind() != tch::Kind::Float && y0.kind() != tch::Kind::BFloat16 && y0.kind() != tch::Kind::Half {
37            return Err(anyhow!("y0 is of unsupported kind {:?}", y0.kind()));
38        }
39        if x_span.kind() != y0.kind() {
40            return Err(anyhow!("x_span and y0 must be of the same kind. Kind of x_span is {:?}. Kind of y0 is {:?}", x_span.kind(), y0.kind()));
41        }
42
43        match self {
44            Self::Euler { step } => solve_euler(f, x_span, y0, *step),
45            Self::RK4 { step } => solve_rk4(f, x_span, y0, *step),
46            Self::ImplicitEuler { step, optimizer } => solve_implicit_euler(f, x_span, y0, *step, optimizer.as_ref()),
47            Self::GLRK4 { step, optimizer } => solve_glrk4(f, x_span, y0, *step, optimizer.as_ref()),
48            Self::RKF45 { rtol, atol, min_step, safety_factor } => solve_rkf45(f, x_span, y0, *rtol, *atol, *min_step, *safety_factor),
49            Self::ROW1 { step } => solve_row1(f, x_span, y0, *step)
50        }
51    }
52}
53
54/// Solves ODE using Euler method
55fn solve_euler(
56    f: tch::CModule,
57    x_span: Tensor,
58    y0: Tensor,
59    step: f64,
60) -> anyhow::Result<(Tensor, Tensor)> {
61    let x_start = x_span.i(0);
62    let x_end = x_span.i(1);
63
64    let mut x = x_start.unsqueeze(0);
65    let mut y = y0.unsqueeze(0);
66
67    let mut all_x = vec![x.copy()];
68    let mut all_y = vec![y.copy()];
69
70    let mut current_step = step;
71    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
72        let remaining = &x_end - &x.squeeze();
73        if remaining.double_value(&[]) < current_step {
74            current_step = remaining.double_value(&[]);
75        }
76
77        let dy = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
78        let dy_rank = dy.size().len();
79        if dy_rank != 1 {
80            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", dy_rank);
81        }
82        
83        y = &y + current_step * &dy;
84        x = &x + current_step;
85
86        all_x.push(x.copy());
87        all_y.push(y.copy());
88    }
89
90    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
91}
92
93/// Solves ODE using Runge-Kutta 4th order method
94fn solve_rk4(
95    f: tch::CModule,
96    x_span: Tensor,
97    y0: Tensor,
98    step: f64,
99) -> anyhow::Result<(Tensor, Tensor)> {
100    let x_start = x_span.i(0);
101    let x_end = x_span.i(1);
102
103    let mut x = x_start.unsqueeze(0);
104    let mut y = y0.unsqueeze(0);
105
106    let mut all_x = vec![x.copy()];
107    let mut all_y = vec![y.copy()];
108
109    let mut current_step = step;
110    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
111        let remaining = &x_end - &x.squeeze();
112        if remaining.double_value(&[]) < current_step {
113            current_step = remaining.double_value(&[]);
114        }
115
116        let k1 = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
117        let k1_rank = k1.size().len();
118        if k1_rank != 1 {
119            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k1_rank);
120        }
121
122        let x_half: Tensor = &x + 0.5 * current_step;
123        let y_half: Tensor = &y + 0.5 * current_step * &k1;
124        let k2 = f.forward_ts(&[x_half.squeeze(), y_half.squeeze()])?;
125        let k2_rank = k2.size().len();
126        if k2_rank != 1 {
127            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k2_rank);
128        }
129
130        let x_half_again: Tensor = &x + 0.5 * current_step;
131        let y_half_again: Tensor = &y + 0.5 * current_step * &k2;
132        let k3 = f.forward_ts(&[x_half_again.squeeze(), y_half_again.squeeze()])?;
133        let k3_rank = k3.size().len();
134        if k3_rank != 1 {
135            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k3_rank);
136        }
137
138        let x_full = &x + current_step;
139        let y_full = &y + current_step * &k3;
140        let k4 = f.forward_ts(&[x_full.squeeze(), y_full.squeeze()])?;
141        let k4_rank = k4.size().len();
142        if k4_rank != 1 {
143            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k4_rank);
144        }
145
146        let step_div_6 = current_step / 6.0;
147        let y_next = &y + step_div_6 * (&k1 + 2.0 * &k2 + 2.0 * &k3 + &k4);
148
149        x = &x + current_step;
150        y = y_next;
151
152        all_x.push(x.copy());
153        all_y.push(y.copy());
154    }
155
156    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
157}
158
159/// Solves ODE using Implicit Euler method with gradient descent optimization
160fn solve_implicit_euler(
161    f: tch::CModule,
162    x_span: Tensor,
163    y0: Tensor,
164    step: f64,
165    optimizer: &dyn optimizers::Optimizer,
166) -> anyhow::Result<(Tensor, Tensor)> {
167    let x_start = x_span.i(0);
168    let x_end = x_span.i(1);
169
170    let mut x = x_start.unsqueeze(0);
171    let mut y = y0.unsqueeze(0);
172
173    let mut all_x = vec![x.copy()];
174    let mut all_y = vec![y.copy()];
175
176    let mut current_step = step;
177    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
178        let remaining = &x_end - &x.squeeze();
179        if remaining.double_value(&[]) < current_step {
180            current_step = remaining.double_value(&[]);
181        }
182
183        let x_next = &x + current_step;
184        let y_prev = y.copy();
185
186        let y_next = optimizer.optimize(
187            &|y_next: &Tensor| {
188                let f_next = f
189                    .forward_ts(&[x_next.squeeze().copy(), y_next.squeeze().copy()])
190                    .unwrap();
191                let y_pred = &y_prev.squeeze() + current_step * &f_next;
192                (y_next - &y_pred).pow_tensor_scalar(2).sum(y_next.kind())
193            },
194            &(&y_prev.detach().squeeze()
195                + current_step * f.forward_ts(&[&x.squeeze(), &y_prev.squeeze()])?),
196        ).map_err( |err| {
197            anyhow!(format!("Optimizer failed with: {}", err))
198        })?;
199
200        y = y_next.unsqueeze(0);
201        x = x_next.copy();
202
203        all_x.push(x.copy());
204        all_y.push(y.copy());
205    }
206
207    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
208}
209
210/// Solves ODE using Gauss-Legendre-Runge-Kutta 4th order method
211fn solve_glrk4(
212    f: tch::CModule,
213    x_span: Tensor,
214    y0: Tensor,
215    step: f64,
216    optimizer: &dyn optimizers::Optimizer,
217) -> anyhow::Result<(Tensor, Tensor)> {
218    let x_start = x_span.i(0);
219    let x_end = x_span.i(1);
220
221    let mut x = x_start.unsqueeze(0);
222    let mut y = y0.unsqueeze(0);
223
224    let mut all_x = vec![x.copy()];
225    let mut all_y = vec![y.copy()];
226
227    let mut current_step = step;
228    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
229        let remaining = &x_end - &x.squeeze();
230        if remaining.double_value(&[]) < current_step {
231            current_step = remaining.double_value(&[]);
232        }
233
234        let k = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
235        let k_rank = k.size().len();
236        if k_rank != 1 {
237            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k_rank);
238        }
239
240        const C1: f64 = 0.2113248654f64;
241        const C2: f64 = 0.7886751346f64;
242        const A11: f64 = 0.25;
243        const A12: f64 = -0.03867513459f64;
244        const A21: f64 = 0.5386751346f64;
245        const A22: f64 = 0.25;
246
247        let first_k1k2_guess = Tensor::cat(
248            &[
249                f.forward_ts(&[
250                    &x.squeeze() + C1 * current_step,
251                    &y.squeeze() + C1 * current_step * &k,
252                ])?,
253                f.forward_ts(&[
254                    &x.squeeze() + C2 * current_step,
255                    &y.squeeze() + C2 * current_step * &k,
256                ])?,
257            ],
258            0,
259        );
260        let k1k2 = optimizer.optimize(
261            &|k1k2_guess| {
262                let diff1 = k1k2_guess.i(0..=1)
263                    - f.forward_ts(&[
264                        &x.squeeze() + C1 * current_step,
265                        &y.squeeze()
266                            + (A11 * k1k2_guess.i(0..=1) + A12 * k1k2_guess.i(2..=3))
267                                * current_step,
268                    ])
269                    .unwrap();
270                let diff2 = k1k2_guess.i(2..=3)
271                    - f.forward_ts(&[
272                        &x.squeeze() + C2 * current_step,
273                        &y.squeeze()
274                            + (A21 * k1k2_guess.i(0..=1) + A22 * k1k2_guess.i(2..=3))
275                                * current_step,
276                    ])
277                    .unwrap();
278
279                diff1.dot(&diff1) + diff2.dot(&diff2)
280            },
281            &first_k1k2_guess,
282        ).map_err( |err| {
283            anyhow!(format!("Optimizer failed with: {}", err))
284        })?;
285        assert!(k1k2.size().len() == 1);
286        assert!(k1k2.size()[0] == 4);
287
288        x = &x + current_step;
289        y = &y + current_step * (0.5 * k1k2.i(0..=1) + 0.5 * k1k2.i(2..=3));
290
291        all_x.push(x.copy());
292        all_y.push(y.copy());
293    }
294
295    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
296}
297
298/// Solves ODE using Runge-Kutta-Fehlberg 45 adaptive method
299fn solve_rkf45(
300    f: tch::CModule,
301    x_span: Tensor,
302    y0: Tensor,
303    rtol: f64,
304    atol: f64,
305    min_step: f64,
306    safety_factor: f64,
307) -> anyhow::Result<(Tensor, Tensor)> {
308    let x_start = x_span.i(0);
309    let x_end = x_span.i(1);
310
311    let mut x = x_start.unsqueeze(0);
312    let mut y = y0.unsqueeze(0);
313
314    let mut all_x = vec![x.copy()];
315    let mut all_y = vec![y.copy()];
316
317    let mut step = (&x_end - &x_start) * 0.1;
318    let safety_factor_tensor = Tensor::from(safety_factor);
319
320    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
321        let remaining = &x_end - &x.squeeze();
322        if remaining.lt_tensor(&step) == Tensor::from(true) {
323            step = remaining.copy();
324        }
325
326        let k1 = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
327        let k1_rank = k1.size().len();
328        if k1_rank != 1 {
329            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k1_rank);
330        }
331
332        let k2 = {
333            let x_step: Tensor = &x + 0.25 * &step;
334            let y_step: Tensor = &y + 0.25 * &step * &k1;
335            let k2_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
336            let k2_rank = k2_unchecked.size().len();
337            if k2_rank != 1 {
338                anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k2_rank);
339            }
340
341            k2_unchecked
342        };
343
344        let k3 = {
345            let x_step: Tensor = &x + 0.375 * &step;
346            let y_step: Tensor = &y + (0.09375 * &step * &k1) + (0.28125 * &step * &k2);
347            let k3_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
348            let k3_rank = k3_unchecked.size().len();
349            if k3_rank != 1 {
350                anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k3_rank);
351            }
352
353            k3_unchecked
354        };
355
356        let k4 = {
357            let x_step: Tensor = &x + (12.0 / 13.0) * &step;
358            let y_step: Tensor = &y
359                + (1932.0 / 2197.0 * &step * &k1)
360                + (-7200.0 / 2197.0 * &step * &k2)
361                + (7296.0 / 2197.0 * &step * &k3);
362            let k4_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
363            let k4_rank = k4_unchecked.size().len();
364            if k4_rank != 1 {
365                anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k4_rank);
366            }
367
368            k4_unchecked
369        };
370
371        let k5 = {
372            let x_step: Tensor = &x + &step;
373            let y_step: Tensor = &y
374                + (439.0 / 216.0 * &step * &k1)
375                + (-8.0 * &step * &k2)
376                + (3680.0 / 513.0 * &step * &k3)
377                + (-845.0 / 4104.0 * &step * &k4);
378            let k5_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
379            let k5_rank = k5_unchecked.size().len();
380            if k5_rank != 1 {
381                anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k5_rank);
382            }
383
384            k5_unchecked
385        };
386
387        let k6 = {
388            let x_step: Tensor = &x + 0.5 * &step;
389            let y_step: Tensor = &y
390                + (-8.0 / 27.0 * &step * &k1)
391                + (2.0 * &step * &k2)
392                + (-3544.0 / 2565.0 * &step * &k3)
393                + (1859.0 / 4104.0 * &step * &k4)
394                + (-11.0 / 40.0 * &step * &k5);
395            let k6_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
396            let k6_rank = k6_unchecked.size().len();
397            if k6_rank != 1 {
398                anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k6_rank);
399            }
400
401            k6_unchecked
402        };
403
404        let next_y4: Tensor = &y
405            + &step
406                * ((25.0 / 216.0 * &k1)
407                    + (1408.0 / 2565.0 * &k3)
408                    + (2197.0 / 4104.0 * &k4)
409                    + (-1.0 / 5.0 * &k5));
410        let next_y5: Tensor = &y
411            + &step
412                * ((16.0 / 135.0 * &k1)
413                    + (6656.0 / 12825.0 * &k3)
414                    + (28561.0 / 56430.0 * &k4)
415                    + (-9.0 / 50.0 * &k5)
416                    + (2.0 / 55.0 * &k6));
417
418        let d = (&next_y4 - &next_y5).abs();
419        let e = next_y5.abs() * rtol + atol;
420
421        let alpha_tensor = (e / d).sqrt().min();
422        let condition = &safety_factor_tensor * &alpha_tensor;
423
424        let condition_met = condition.lt(1.0);
425        let condition_met_bool: bool = condition_met == Tensor::from(true);
426
427        if condition_met_bool {
428            step = &step * &condition;
429            if step.double_value(&[]) < min_step {
430                return Err(anyhow!("Required step is smaller than minimal step"));
431            }
432        } else {
433            y = next_y4;
434            x = &x + &step;
435            all_x.push(x.copy());
436            all_y.push(y.copy());
437
438            let new_step = &step * &condition;
439            let max_step = &step * 5.0;
440            step = new_step.fmin(&max_step);
441        }
442    }
443
444    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
445}
446
447/// Solves ODE using first-order Rosenbrock method (Row1)
448fn solve_row1(
449    f: tch::CModule,
450    x_span: Tensor,
451    y0: Tensor,
452    step: f64,
453) -> anyhow::Result<(Tensor, Tensor)> {
454    let x_start = x_span.i(0);
455    let x_end = x_span.i(1);
456
457    let mut x = x_start.unsqueeze(0);
458    let mut y = y0.unsqueeze(0);
459
460    let mut all_x = vec![x.copy()];
461    let mut all_y = vec![y.copy()];
462
463    while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
464        let remaining = &x_end - &x.squeeze();
465        let mut current_step = step;
466        if remaining.double_value(&[]) < step {
467            current_step = remaining.double_value(&[]);
468        }
469
470        let x_prev = x.copy();
471        let y_prev = y.copy().squeeze();
472
473        let jacobian = compute_jacobian(
474            |y| {
475                f.forward_ts(&[x_prev.squeeze().copy(), y.copy()])
476                    .unwrap()
477                    .squeeze()
478            },
479            &y_prev,
480        );
481        let f_current = f
482            .forward_ts(&[x_prev.squeeze().copy(), y_prev.copy()])?;
483        let f_current_rank = f_current.size().len();
484        if f_current_rank != 1 {
485            anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", f_current_rank);
486        }
487
488        let n = jacobian.size()[0];
489        let eye = Tensor::eye(n, (tch::Kind::Float, jacobian.device()));
490        let step_j = current_step * &jacobian;
491        let inv_matrix = (eye - step_j).inverse();
492
493        let delta_y = inv_matrix.matmul(&f_current);
494        let y_next = y_prev + current_step * delta_y;
495
496        x = &x_prev + current_step;
497        y = y_next.unsqueeze(0);
498
499        all_x.push(x.copy());
500        all_y.push(y.copy());
501    }
502
503    Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
504}
505
506/// Computes the Jacobian matrix of a function f at point x
507fn compute_jacobian<F>(f: F, x: &Tensor) -> Tensor
508where
509    F: Fn(&Tensor) -> Tensor,
510{
511    assert_eq!(x.dim(), 1, "x must be 1-dimensional");
512    let mut x_with_grad = x.detach().copy().set_requires_grad(true);
513    let y = f(&x_with_grad);
514    assert_eq!(y.dim(), 1, "y must be 1-dimensional");
515
516    let y_size = y.size()[0];
517    let mut grads = Vec::new();
518
519    for i in 0..y_size {
520        let yi = y.i(i);
521        //yi.backward();
522        //let grad = x_with_grad.grad().copy();
523        let grad = Tensor::run_backward(&[yi], &[&x_with_grad], true, false)[0].copy();
524        grads.push(grad.unsqueeze(0));
525        x_with_grad.zero_grad();
526    }
527
528    Tensor::cat(&grads, 0)
529}