Skip to main content

solver/numeric/
ode.rs

1pub trait ODESolver {
2    /// Solves an ODE system dy/dt = f(t, y) from t_start to t_end.
3    ///
4    /// # Arguments
5    /// * `f` - The derivative function f(t, y).
6    /// * `y0` - Initial state vector.
7    /// * `t_span` - Tuple (t_start, t_end).
8    /// * `dt` - Step size.
9    fn solve<F>(f: F, y0: Vec<f64>, t_span: (f64, f64), dt: f64) -> (Vec<f64>, Vec<Vec<f64>>)
10    where
11        F: FnMut(f64, &Vec<f64>) -> Vec<f64>;
12}
13
14/// Euler Method O(dt)
15pub struct Euler;
16
17impl ODESolver for Euler {
18    fn solve<F>(mut f: F, y0: Vec<f64>, t_span: (f64, f64), dt: f64) -> (Vec<f64>, Vec<Vec<f64>>)
19    where
20        F: FnMut(f64, &Vec<f64>) -> Vec<f64>,
21    {
22        let (t_start, t_end) = t_span;
23        let n_steps = ((t_end - t_start).abs() / dt).ceil() as usize;
24
25        let mut t_values = Vec::with_capacity(n_steps + 1);
26        let mut y_values = Vec::with_capacity(n_steps + 1);
27
28        let mut t = t_start;
29        let mut y = y0;
30
31        t_values.push(t);
32        y_values.push(y.clone());
33
34        for _ in 0..n_steps {
35            let dy = f(t, &y);
36            y = add(&y, &scale(&dy, dt));
37            t += dt;
38
39            t_values.push(t);
40            y_values.push(y.clone());
41        }
42
43        (t_values, y_values)
44    }
45}
46
47/// Runge-Kutta 4th Order Method O(dt^4)
48pub struct RungeKutta4;
49
50impl ODESolver for RungeKutta4 {
51    fn solve<F>(mut f: F, y0: Vec<f64>, t_span: (f64, f64), dt: f64) -> (Vec<f64>, Vec<Vec<f64>>)
52    where
53        F: FnMut(f64, &Vec<f64>) -> Vec<f64>,
54    {
55        let (t_start, t_end) = t_span;
56        let n_steps = ((t_end - t_start).abs() / dt).ceil() as usize;
57
58        let mut t_values = Vec::with_capacity(n_steps + 1);
59        let mut y_values = Vec::with_capacity(n_steps + 1);
60
61        let mut t = t_start;
62        let mut y = y0;
63
64        t_values.push(t);
65        y_values.push(y.clone());
66
67        for _ in 0..n_steps {
68            let k1 = f(t, &y);
69            let k2 = f(t + 0.5 * dt, &add(&y, &scale(&k1, 0.5 * dt)));
70            let k3 = f(t + 0.5 * dt, &add(&y, &scale(&k2, 0.5 * dt)));
71            let k4 = f(t + dt, &add(&y, &scale(&k3, dt)));
72
73            // y_new = y + (dt/6) * (k1 + 2*k2 + 2*k3 + k4)
74            let delta = scale(
75                &add(&add(&k1, &scale(&k2, 2.0)), &add(&scale(&k3, 2.0), &k4)),
76                dt / 6.0,
77            );
78            y = add(&y, &delta);
79            t += dt;
80
81            t_values.push(t);
82            y_values.push(y.clone());
83        }
84
85        (t_values, y_values)
86    }
87}
88
89// Helper functions for vector arithmetic
90fn add(a: &[f64], b: &[f64]) -> Vec<f64> {
91    a.iter().zip(b.iter()).map(|(x, y)| x + y).collect()
92}
93
94fn scale(a: &[f64], scalar: f64) -> Vec<f64> {
95    a.iter().map(|x| x * scalar).collect()
96}
97
98use crate::linalg::eigen::eigen_symmetric_tridiagonal;
99
100/// Solver for Linear Time-Invariant (LTI) systems y' = My using spectral decomposition.
101/// Efficient for constant matrices, as it avoids time-stepping.
102pub struct LinearSpectralSolver;
103
104impl LinearSpectralSolver {
105    /// Solves y(t) = exp(A * t) * y0 for a symmetric tridiagonal matrix A.
106    ///
107    /// # Arguments
108    /// * `diag` - Diagonal elements of A.
109    /// * `off_diag` - Off-diagonal elements of A.
110    /// * `y0` - Initial state vector at t=0.
111    /// * `t` - Time horizon to propagate to.
112    pub fn solve_tridiagonal(diag: Vec<f64>, off_diag: Vec<f64>, y0: &[f64], t: f64) -> Vec<f64> {
113        let n = diag.len();
114
115        // 1. Diagonalize: A = Q * Lambda * Q^T
116        let (evals, evecs) = eigen_symmetric_tridiagonal(diag, off_diag);
117
118        // 2. Project y0 onto basis: c = Q^T * y0
119        // coeffs[k] = dot(evec_k, y0)
120        let mut coeffs = vec![0.0; n];
121        for k in 0..n {
122            let mut dot = 0.0;
123            for i in 0..n {
124                // evecs are stored column-major or flattened?
125                // evecs[i + k*n] is i-th component of k-th eigenvector
126                dot += evecs[i + k * n] * y0[i];
127            }
128            coeffs[k] = dot;
129        }
130
131        // 3. Evolve in eigen-basis and reconstruct:
132        // y(t) = Q * (exp(lambda * t) .* c)
133
134        // Numerically stable evolution:
135        // Find max_exponent = max(lambda_k * t)
136        // Shift exponents: exp(lambda_k * t) = exp((lambda_k * t) - max_exponent) * exp(max_exponent)
137        // This ensures all computed exponentials are <= 1.0, preventing overflow.
138        // The result is scaled by exp(-max_exponent) relative to the true value, which avoids
139        // underflow to zero when all values are extremely small (common in diffusion with absorption).
140        // Since we are typically interested in ratios of y components (for spreads), this scaling cancels out.
141
142        let mut max_exp = f64::NEG_INFINITY;
143        for &lambda in &evals {
144            let val = lambda * t;
145            if val > max_exp {
146                max_exp = val;
147            }
148        }
149
150        let mut y_t = vec![0.0; n];
151        for i in 0..n {
152            let mut sum = 0.0;
153            for k in 0..n {
154                // Combine: coeff_k * exp(lambda_k * t - max_exp) * evec_k[i]
155                sum += coeffs[k] * (evals[k] * t - max_exp).exp() * evecs[i + k * n];
156            }
157            y_t[i] = sum;
158        }
159
160        y_t
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    #[test]
169    fn test_ode_solver_exponential_decay() {
170        // dy/dt = -y, y(0) = 1.0
171        // Analytic solution: y(t) = exp(-t)
172
173        let f = |_t: f64, y: &Vec<f64>| -> Vec<f64> { vec![-y[0]] };
174
175        let y0 = vec![1.0];
176        let t_span = (0.0, 1.0);
177        let dt = 0.01;
178
179        let (_, y) = RungeKutta4::solve(f, y0, t_span, dt);
180
181        let final_y = y.last().unwrap()[0];
182        let expected = (-1.0_f64).exp();
183
184        assert!(
185            (final_y - expected).abs() < 1e-4,
186            "RK4 approximation should be close to exact solution"
187        );
188    }
189
190    #[test]
191    fn test_ode_solver_system_harmonic_oscillator() {
192        // Harmonic oscillator: y'' = -y
193        // let y1 = y, y2 = y'
194        // y1' = y2
195        // y2' = -y1
196        // y(0) = [0, 1] (starts at 0 with velocity 1) -> sin(t)
197
198        let f = |_t: f64, y: &Vec<f64>| -> Vec<f64> {
199            let y1 = y[0];
200            let y2 = y[1];
201            vec![y2, -y1]
202        };
203
204        let y0 = vec![0.0, 1.0];
205        let t_span = (0.0, std::f64::consts::PI / 2.0); // Integrate to pi/2
206        let dt = 0.01;
207
208        let (_, y) = RungeKutta4::solve(f, y0, t_span, dt);
209
210        let final_y1 = y.last().unwrap()[0];
211        let expected = 1.0; // sin(pi/2) = 1
212
213        assert!(
214            (final_y1 - expected).abs() < 1e-4,
215            "RK4 Harmonic Oscillator"
216        );
217    }
218}