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
98#[cfg(test)]
99mod tests {
100    use super::*;
101
102    #[test]
103    fn test_ode_solver_exponential_decay() {
104        // dy/dt = -y, y(0) = 1.0
105        // Analytic solution: y(t) = exp(-t)
106
107        let f = |_t: f64, y: &Vec<f64>| -> Vec<f64> { vec![-y[0]] };
108
109        let y0 = vec![1.0];
110        let t_span = (0.0, 1.0);
111        let dt = 0.01;
112
113        let (_, y) = RungeKutta4::solve(f, y0, t_span, dt);
114
115        let final_y = y.last().unwrap()[0];
116        let expected = (-1.0_f64).exp();
117
118        assert!(
119            (final_y - expected).abs() < 1e-4,
120            "RK4 approximation should be close to exact solution"
121        );
122    }
123
124    #[test]
125    fn test_ode_solver_system_harmonic_oscillator() {
126        // Harmonic oscillator: y'' = -y
127        // let y1 = y, y2 = y'
128        // y1' = y2
129        // y2' = -y1
130        // y(0) = [0, 1] (starts at 0 with velocity 1) -> sin(t)
131
132        let f = |_t: f64, y: &Vec<f64>| -> Vec<f64> {
133            let y1 = y[0];
134            let y2 = y[1];
135            vec![y2, -y1]
136        };
137
138        let y0 = vec![0.0, 1.0];
139        let t_span = (0.0, std::f64::consts::PI / 2.0); // Integrate to pi/2
140        let dt = 0.01;
141
142        let (_, y) = RungeKutta4::solve(f, y0, t_span, dt);
143
144        let final_y1 = y.last().unwrap()[0];
145        let expected = 1.0; // sin(pi/2) = 1
146
147        assert!(
148            (final_y1 - expected).abs() < 1e-4,
149            "RK4 Harmonic Oscillator"
150        );
151    }
152}
153
154use crate::linalg::eigen::eigen_symmetric_tridiagonal;
155
156/// Solver for Linear Time-Invariant (LTI) systems y' = My using spectral decomposition.
157/// Efficient for constant matrices, as it avoids time-stepping.
158pub struct LinearSpectralSolver;
159
160impl LinearSpectralSolver {
161    /// Solves y(t) = exp(A * t) * y0 for a symmetric tridiagonal matrix A.
162    ///
163    /// # Arguments
164    /// * `diag` - Diagonal elements of A.
165    /// * `off_diag` - Off-diagonal elements of A.
166    /// * `y0` - Initial state vector at t=0.
167    /// * `t` - Time horizon to propagate to.
168    pub fn solve_tridiagonal(diag: Vec<f64>, off_diag: Vec<f64>, y0: &[f64], t: f64) -> Vec<f64> {
169        let n = diag.len();
170
171        // 1. Diagonalize: A = Q * Lambda * Q^T
172        let (evals, evecs) = eigen_symmetric_tridiagonal(diag, off_diag);
173
174        // 2. Project y0 onto basis: c = Q^T * y0
175        // coeffs[k] = dot(evec_k, y0)
176        let mut coeffs = vec![0.0; n];
177        for k in 0..n {
178            let mut dot = 0.0;
179            for i in 0..n {
180                // evecs are stored column-major or flattened?
181                // evecs[i + k*n] is i-th component of k-th eigenvector
182                dot += evecs[i + k * n] * y0[i];
183            }
184            coeffs[k] = dot;
185        }
186
187        // 3. Evolve in eigen-basis and reconstruct:
188        // y(t) = Q * (exp(lambda * t) .* c)
189
190        // Numerically stable evolution:
191        // Find max_exponent = max(lambda_k * t)
192        // Shift exponents: exp(lambda_k * t) = exp((lambda_k * t) - max_exponent) * exp(max_exponent)
193        // This ensures all computed exponentials are <= 1.0, preventing overflow.
194        // The result is scaled by exp(-max_exponent) relative to the true value, which avoids
195        // underflow to zero when all values are extremely small (common in diffusion with absorption).
196        // Since we are typically interested in ratios of y components (for spreads), this scaling cancels out.
197
198        let mut max_exp = f64::NEG_INFINITY;
199        for &lambda in &evals {
200            let val = lambda * t;
201            if val > max_exp {
202                max_exp = val;
203            }
204        }
205
206        let mut y_t = vec![0.0; n];
207        for i in 0..n {
208            let mut sum = 0.0;
209            for k in 0..n {
210                // Combine: coeff_k * exp(lambda_k * t - max_exp) * evec_k[i]
211                sum += coeffs[k] * (evals[k] * t - max_exp).exp() * evecs[i + k * n];
212            }
213            y_t[i] = sum;
214        }
215
216        y_t
217    }
218}