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(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
20/// Validates that a tensor contains only finite values.
21/// Returns an error if any NaN or Inf values are detected.
22fn 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
33/// Validates that a scalar value is finite.
34fn 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        // Validate x_span
78        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        // Validate solver parameters
86        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 y0 - check it's finite
137        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        // Validate function f
157        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 derivative output is finite
196        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
281/// Solves ODE using Euler method
282fn 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        // Compute next state
332        y = &y + current_step * &dy;
333
334        // Critical: validate new state is finite before proceeding
335        validate_finite_tensor(&y, "state after Euler update (NaN/Inf propagating)")?;
336
337        x = &x + current_step;
338
339        // Validate x remains finite
340        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
375/// Solves ODE using Runge-Kutta 4th order method
376fn 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        // Stage k1
407        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        // Stage k2
416        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        // Stage k3
429        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        // Stage k4
445        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        // Compute next state using weighted average of stages
458        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        // Critical validation after full RK4 step
462        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
499/// Solves ODE using Implicit Euler method with gradient descent optimization
500fn 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        // Create derivative function for current x
535        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        // Initial guess based on explicit Euler
547        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        // Run optimizer (may fail gracefully internally)
552        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        // Critical: validate optimizer output before accepting
557        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
597/// Solves ODE using Gauss-Legendre-Runge-Kutta 4th order method
598fn 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        // Initial guess for k1, k2
646        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        // Define loss function for optimization
665        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        // Run optimizer
690        let k1k2 = optimizer
691            .optimize(&loss_fn, &first_k1k2_guess)
692            .map_err(|err| anyhow!(format!("GLRK4 optimizer failed with: {}", err)))?;
693
694        // Validate optimizer output
695        validate_finite_tensor(
696            &k1k2,
697            "GLRK4 stage coefficients after optimization (NaN/Inf)",
698        )?;
699
700        // Compute final state update
701        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 final state
707        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
741/// One RKF45 step
742fn 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    // Stage k1
752    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    // Stage k2
761    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    // Stage k3
774    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    // Stage k4
787    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    // Stage k5
803    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    // Stage k6
820    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    // Compute 4th and 5th order solutions
838    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
854/// Solves ODE using Runge-Kutta-Fehlberg 45 adaptive method
855fn 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        // Compute error estimate
892        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        // Debug
899        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        // Compute step size adjustment
909        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            // Step rejected - shrink and retry
918            consecutive_rejections += 1;
919            total_rejections += 1;
920
921            // Warning for consecutive rejections
922            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            // Warning for many total rejections
931            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            // Warning for very small step approaching min_step
940            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            // Warning for step below min_step
952            if step < min_step {
953                return Err(anyhow!("Required step is smaller than minimal step"));
954            }
955        } else {
956            // Accept the step
957            consecutive_rejections = 0;
958            total_accepted += 1;
959
960            // At last step, special handling
961            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 accepted state
976            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    // Final efficiency summary
988    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
1003/// Solves ODE using first-order Rosenbrock method (Row1)
1004fn 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        // Compute Jacobian
1040        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 Jacobian is finite
1052        validate_finite_tensor(&jacobian, "Jacobian matrix in ROW1 (NaN/Inf)")?;
1053
1054        // Evaluate function at current point
1055        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        // Compute (I - h*J)^(-1) * f
1078        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        // Warn about ill-conditioning before inversion (one-shot)
1084        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        // Warn about inverse magnitude (one-shot)
1098        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        // Critical validation after ROW1 update
1113        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
1152/// Computes the Jacobian matrix of a function f at point x
1153fn 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}