use nalgebra::{DMatrix, DVector};
pub struct Sindy {
pub names: Vec<String>,
pub coeffs: DMatrix<f64>,
}
pub fn monomial_exponents(n_vars: usize, degree: usize) -> Vec<Vec<usize>> {
let mut out = Vec::new();
fn rec(pos: usize, n: usize, left: usize, cur: &mut Vec<usize>, out: &mut Vec<Vec<usize>>) {
if pos == n {
out.push(cur.clone());
return;
}
for e in 0..=left {
cur.push(e);
rec(pos + 1, n, left - e, cur, out);
cur.pop();
}
}
rec(0, n_vars, degree, &mut Vec::new(), &mut out);
out.sort_by_key(|v| (v.iter().sum::<usize>(), v.clone()));
out
}
fn term_name(exps: &[usize]) -> String {
let vars = ["x", "y", "z", "w"];
let parts: Vec<String> = exps
.iter()
.enumerate()
.filter(|(_, e)| **e > 0)
.map(|(i, &e)| if e == 1 { vars[i].to_string() } else { format!("{}^{}", vars[i], e) })
.collect();
if parts.is_empty() { "1".to_string() } else { parts.join(" ") }
}
fn eval_term(exps: &[usize], state: &[f64]) -> f64 {
exps.iter().zip(state.iter()).map(|(&e, &v)| v.powi(e as i32)).product()
}
fn lstsq(a: &DMatrix<f64>, b: &DVector<f64>) -> DVector<f64> {
a.clone().svd(true, true).solve(b, 1e-12).unwrap_or_else(|_| DVector::zeros(a.ncols()))
}
fn stlsq(theta: &DMatrix<f64>, d: &DVector<f64>, lambda: f64, iters: usize) -> DVector<f64> {
let n = theta.ncols();
let mut xi = lstsq(theta, d);
for _ in 0..iters {
let active: Vec<usize> = (0..n).filter(|&j| xi[j].abs() >= lambda).collect();
for j in 0..n {
if xi[j].abs() < lambda {
xi[j] = 0.0;
}
}
if active.is_empty() {
break;
}
let sub = DMatrix::from_columns(&active.iter().map(|&j| theta.column(j).into_owned()).collect::<Vec<_>>());
let sol = lstsq(&sub, d);
for (k, &j) in active.iter().enumerate() {
xi[j] = sol[k];
}
}
xi
}
impl Sindy {
pub fn fit(states: &[Vec<f64>], derivs: &[Vec<f64>], degree: usize, lambda: f64) -> Sindy {
let n_vars = states[0].len();
let exps = monomial_exponents(n_vars, degree);
let n_feat = exps.len();
let n_samp = states.len();
let mut theta = DMatrix::zeros(n_samp, n_feat);
for (i, s) in states.iter().enumerate() {
for (j, e) in exps.iter().enumerate() {
theta[(i, j)] = eval_term(e, s);
}
}
let mut coeffs = DMatrix::zeros(n_feat, n_vars);
for dim in 0..n_vars {
let d = DVector::from_iterator(n_samp, derivs.iter().map(|dv| dv[dim]));
let xi = stlsq(&theta, &d, lambda, 12);
for j in 0..n_feat {
coeffs[(j, dim)] = xi[j];
}
}
Sindy { names: exps.iter().map(|e| term_name(e)).collect(), coeffs }
}
pub fn n_active(&self) -> usize {
self.coeffs.iter().filter(|&&c| c != 0.0).count()
}
pub fn equation(&self, dim: usize) -> String {
let lhs = ["ẋ", "ẏ", "ż", "ẇ"][dim.min(3)];
let terms: Vec<String> = (0..self.names.len())
.filter(|&j| self.coeffs[(j, dim)] != 0.0)
.map(|j| {
let c = self.coeffs[(j, dim)];
if self.names[j] == "1" { format!("{c:.3}") } else { format!("{c:.3} {}", self.names[j]) }
})
.collect();
format!("{lhs} = {}", if terms.is_empty() { "0".to_string() } else { terms.join(" + ") })
}
}
#[cfg(test)]
mod tests {
use super::*;
fn duffing(x: f64, y: f64) -> (f64, f64) {
(y, -x - 0.3 * x * x * x - 0.1 * y)
}
fn trajectory(x0: f64, y0: f64, n: usize, dt: f64) -> (Vec<Vec<f64>>, Vec<Vec<f64>>) {
let (mut x, mut y) = (x0, y0);
let (mut states, mut derivs) = (vec![], vec![]);
for _ in 0..n {
let (dx, dy) = duffing(x, y);
states.push(vec![x, y]);
derivs.push(vec![dx, dy]); let (k1x, k1y) = duffing(x, y);
let (k2x, k2y) = duffing(x + 0.5 * dt * k1x, y + 0.5 * dt * k1y);
let (k3x, k3y) = duffing(x + 0.5 * dt * k2x, y + 0.5 * dt * k2y);
let (k4x, k4y) = duffing(x + dt * k3x, y + dt * k3y);
x += dt / 6.0 * (k1x + 2.0 * k2x + 2.0 * k3x + k4x);
y += dt / 6.0 * (k1y + 2.0 * k2y + 2.0 * k3y + k4y);
}
(states, derivs)
}
#[test]
fn sindy_discovers_the_duffing_oscillator() {
let (s1, d1) = trajectory(1.5, 0.0, 1200, 0.01);
let (s2, d2) = trajectory(0.4, 1.2, 1200, 0.01); let states: Vec<Vec<f64>> = s1.into_iter().chain(s2).collect();
let derivs: Vec<Vec<f64>> = d1.into_iter().chain(d2).collect();
let sindy = Sindy::fit(&states, &derivs, 3, 0.05);
let idx = |name: &str| sindy.names.iter().position(|n| n == name).unwrap();
assert!((sindy.coeffs[(idx("y"), 0)] - 1.0).abs() < 1e-2, "ẋ should be y: {}", sindy.equation(0));
assert!((sindy.coeffs[(idx("x"), 1)] - (-1.0)).abs() < 2e-2, "x term: {}", sindy.equation(1));
assert!((sindy.coeffs[(idx("y"), 1)] - (-0.1)).abs() < 2e-2, "y term: {}", sindy.equation(1));
assert!((sindy.coeffs[(idx("x^3"), 1)] - (-0.3)).abs() < 2e-2, "x³ term: {}", sindy.equation(1));
assert_eq!(sindy.n_active(), 4, "should recover exactly 4 terms, got {}: {} ; {}", sindy.n_active(), sindy.equation(0), sindy.equation(1));
}
#[test]
fn too_large_a_threshold_oversparsifies() {
let (states, derivs) = trajectory(1.5, 0.0, 1500, 0.01);
let sindy = Sindy::fit(&states, &derivs, 3, 0.5);
assert!(sindy.n_active() < 4, "an over-large threshold should drop real terms");
}
#[test]
fn monomial_library_has_the_expected_terms() {
let e = monomial_exponents(2, 3);
assert_eq!(e.len(), 10, "2 vars, degree 3 → 10 terms");
let names: Vec<String> = e.iter().map(|x| term_name(x)).collect();
assert!(names.contains(&"1".to_string()) && names.contains(&"x^3".to_string()) && names.contains(&"x y".to_string()));
}
}