1pub trait ODESolver {
2 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
14pub 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
47pub 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 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
89fn 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
100pub struct LinearSpectralSolver;
103
104impl LinearSpectralSolver {
105 pub fn solve_tridiagonal(diag: Vec<f64>, off_diag: Vec<f64>, y0: &[f64], t: f64) -> Vec<f64> {
113 let n = diag.len();
114
115 let (evals, evecs) = eigen_symmetric_tridiagonal(diag, off_diag);
117
118 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 dot += evecs[i + k * n] * y0[i];
127 }
128 coeffs[k] = dot;
129 }
130
131 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 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 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 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); 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; assert!(
214 (final_y1 - expected).abs() < 1e-4,
215 "RK4 Harmonic Oscillator"
216 );
217 }
218}