1use nalgebra::{DMatrix, DVector};
16
17pub struct Sindy {
20 pub names: Vec<String>,
21 pub coeffs: DMatrix<f64>,
23}
24
25pub fn monomial_exponents(n_vars: usize, degree: usize) -> Vec<Vec<usize>> {
27 let mut out = Vec::new();
28 fn rec(pos: usize, n: usize, left: usize, cur: &mut Vec<usize>, out: &mut Vec<Vec<usize>>) {
30 if pos == n {
31 out.push(cur.clone());
32 return;
33 }
34 for e in 0..=left {
35 cur.push(e);
36 rec(pos + 1, n, left - e, cur, out);
37 cur.pop();
38 }
39 }
40 rec(0, n_vars, degree, &mut Vec::new(), &mut out);
41 out.sort_by_key(|v| (v.iter().sum::<usize>(), v.clone()));
43 out
44}
45
46fn term_name(exps: &[usize]) -> String {
47 let vars = ["x", "y", "z", "w"];
48 let parts: Vec<String> = exps
49 .iter()
50 .enumerate()
51 .filter(|(_, e)| **e > 0)
52 .map(|(i, &e)| if e == 1 { vars[i].to_string() } else { format!("{}^{}", vars[i], e) })
53 .collect();
54 if parts.is_empty() { "1".to_string() } else { parts.join(" ") }
55}
56
57fn eval_term(exps: &[usize], state: &[f64]) -> f64 {
58 exps.iter().zip(state.iter()).map(|(&e, &v)| v.powi(e as i32)).product()
59}
60
61fn lstsq(a: &DMatrix<f64>, b: &DVector<f64>) -> DVector<f64> {
63 a.clone().svd(true, true).solve(b, 1e-12).unwrap_or_else(|_| DVector::zeros(a.ncols()))
64}
65
66fn stlsq(theta: &DMatrix<f64>, d: &DVector<f64>, lambda: f64, iters: usize) -> DVector<f64> {
68 let n = theta.ncols();
69 let mut xi = lstsq(theta, d);
70 for _ in 0..iters {
71 let active: Vec<usize> = (0..n).filter(|&j| xi[j].abs() >= lambda).collect();
72 for j in 0..n {
74 if xi[j].abs() < lambda {
75 xi[j] = 0.0;
76 }
77 }
78 if active.is_empty() {
79 break;
80 }
81 let sub = DMatrix::from_columns(&active.iter().map(|&j| theta.column(j).into_owned()).collect::<Vec<_>>());
83 let sol = lstsq(&sub, d);
84 for (k, &j) in active.iter().enumerate() {
85 xi[j] = sol[k];
86 }
87 }
88 xi
89}
90
91impl Sindy {
92 pub fn fit(states: &[Vec<f64>], derivs: &[Vec<f64>], degree: usize, lambda: f64) -> Sindy {
95 let n_vars = states[0].len();
96 let exps = monomial_exponents(n_vars, degree);
97 let n_feat = exps.len();
98 let n_samp = states.len();
99
100 let mut theta = DMatrix::zeros(n_samp, n_feat);
102 for (i, s) in states.iter().enumerate() {
103 for (j, e) in exps.iter().enumerate() {
104 theta[(i, j)] = eval_term(e, s);
105 }
106 }
107
108 let mut coeffs = DMatrix::zeros(n_feat, n_vars);
110 for dim in 0..n_vars {
111 let d = DVector::from_iterator(n_samp, derivs.iter().map(|dv| dv[dim]));
112 let xi = stlsq(&theta, &d, lambda, 12);
113 for j in 0..n_feat {
114 coeffs[(j, dim)] = xi[j];
115 }
116 }
117
118 Sindy { names: exps.iter().map(|e| term_name(e)).collect(), coeffs }
119 }
120
121 pub fn n_active(&self) -> usize {
123 self.coeffs.iter().filter(|&&c| c != 0.0).count()
124 }
125
126 pub fn equation(&self, dim: usize) -> String {
128 let lhs = ["ẋ", "ẏ", "ż", "ẇ"][dim.min(3)];
129 let terms: Vec<String> = (0..self.names.len())
130 .filter(|&j| self.coeffs[(j, dim)] != 0.0)
131 .map(|j| {
132 let c = self.coeffs[(j, dim)];
133 if self.names[j] == "1" { format!("{c:.3}") } else { format!("{c:.3} {}", self.names[j]) }
134 })
135 .collect();
136 format!("{lhs} = {}", if terms.is_empty() { "0".to_string() } else { terms.join(" + ") })
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143
144 fn duffing(x: f64, y: f64) -> (f64, f64) {
146 (y, -x - 0.3 * x * x * x - 0.1 * y)
147 }
148
149 fn trajectory(x0: f64, y0: f64, n: usize, dt: f64) -> (Vec<Vec<f64>>, Vec<Vec<f64>>) {
150 let (mut x, mut y) = (x0, y0);
151 let (mut states, mut derivs) = (vec![], vec![]);
152 for _ in 0..n {
153 let (dx, dy) = duffing(x, y);
154 states.push(vec![x, y]);
155 derivs.push(vec![dx, dy]); let (k1x, k1y) = duffing(x, y);
158 let (k2x, k2y) = duffing(x + 0.5 * dt * k1x, y + 0.5 * dt * k1y);
159 let (k3x, k3y) = duffing(x + 0.5 * dt * k2x, y + 0.5 * dt * k2y);
160 let (k4x, k4y) = duffing(x + dt * k3x, y + dt * k3y);
161 x += dt / 6.0 * (k1x + 2.0 * k2x + 2.0 * k3x + k4x);
162 y += dt / 6.0 * (k1y + 2.0 * k2y + 2.0 * k3y + k4y);
163 }
164 (states, derivs)
165 }
166
167 #[test]
168 fn sindy_discovers_the_duffing_oscillator() {
169 let (s1, d1) = trajectory(1.5, 0.0, 1200, 0.01);
172 let (s2, d2) = trajectory(0.4, 1.2, 1200, 0.01); let states: Vec<Vec<f64>> = s1.into_iter().chain(s2).collect();
174 let derivs: Vec<Vec<f64>> = d1.into_iter().chain(d2).collect();
175
176 let sindy = Sindy::fit(&states, &derivs, 3, 0.05);
177 let idx = |name: &str| sindy.names.iter().position(|n| n == name).unwrap();
178
179 assert!((sindy.coeffs[(idx("y"), 0)] - 1.0).abs() < 1e-2, "ẋ should be y: {}", sindy.equation(0));
181 assert!((sindy.coeffs[(idx("x"), 1)] - (-1.0)).abs() < 2e-2, "x term: {}", sindy.equation(1));
183 assert!((sindy.coeffs[(idx("y"), 1)] - (-0.1)).abs() < 2e-2, "y term: {}", sindy.equation(1));
184 assert!((sindy.coeffs[(idx("x^3"), 1)] - (-0.3)).abs() < 2e-2, "x³ term: {}", sindy.equation(1));
185 assert_eq!(sindy.n_active(), 4, "should recover exactly 4 terms, got {}: {} ; {}", sindy.n_active(), sindy.equation(0), sindy.equation(1));
187 }
188
189 #[test]
190 fn too_large_a_threshold_oversparsifies() {
191 let (states, derivs) = trajectory(1.5, 0.0, 1500, 0.01);
193 let sindy = Sindy::fit(&states, &derivs, 3, 0.5);
194 assert!(sindy.n_active() < 4, "an over-large threshold should drop real terms");
195 }
196
197 #[test]
198 fn monomial_library_has_the_expected_terms() {
199 let e = monomial_exponents(2, 3);
200 assert_eq!(e.len(), 10, "2 vars, degree 3 → 10 terms");
201 let names: Vec<String> = e.iter().map(|x| term_name(x)).collect();
202 assert!(names.contains(&"1".to_string()) && names.contains(&"x^3".to_string()) && names.contains(&"x y".to_string()));
203 }
204}