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
98#[cfg(test)]
99mod tests {
100 use super::*;
101
102 #[test]
103 fn test_ode_solver_exponential_decay() {
104 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 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); 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; assert!(
148 (final_y1 - expected).abs() < 1e-4,
149 "RK4 Harmonic Oscillator"
150 );
151 }
152}
153
154use crate::linalg::eigen::eigen_symmetric_tridiagonal;
155
156pub struct LinearSpectralSolver;
159
160impl LinearSpectralSolver {
161 pub fn solve_tridiagonal(diag: Vec<f64>, off_diag: Vec<f64>, y0: &[f64], t: f64) -> Vec<f64> {
169 let n = diag.len();
170
171 let (evals, evecs) = eigen_symmetric_tridiagonal(diag, off_diag);
173
174 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 dot += evecs[i + k * n] * y0[i];
183 }
184 coeffs[k] = dot;
185 }
186
187 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 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}