Skip to main content

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