//! Lie derivatives for structural identifiability analysis. //! //! The Lie derivative of a function h(x) along the vector field f(x,p) is: //! L_f h = Σᵢ (∂h/∂xᵢ) · fᵢ(x,p) //! //! Higher-order Lie derivatives: L_f^k h = L_f(L_f^{k-1} h) use std::sync::Arc; use symclaw_core::ast::Expr; use symclaw_core::differentiate::differentiate; use symclaw_core::interner::Symbol; use symclaw_core::simplify::simplify; /// Compute the Lie derivative of `h` along vector field `f`. /// /// - `h`: the function to differentiate (depends on states) /// - `f`: per-state RHS expressions fᵢ(x,p) — same order as `state_vars` /// - `state_vars`: symbols for the state variables x₁,…,xₙ /// /// Returns: L_f h = Σᵢ (∂h/∂xᵢ) · fᵢ #[must_use] pub fn lie_derivative(h: &Expr, f: &[Arc], state_vars: &[Symbol]) -> Arc { assert_eq!( f.len(), state_vars.len(), "f and state_vars must have same length" ); let mut terms: Vec> = Vec::new(); for (fi, &xi) in f.iter().zip(state_vars.iter()) { let dh_dxi = differentiate(&Arc::new(h.clone()), xi); let term = Expr::mul(vec![dh_dxi, fi.clone()]); let simplified = simplify(&(*term).clone()); // Skip zero terms use num_traits::Zero; if !matches!(simplified.as_ref(), Expr::Num(r) if r.is_zero()) { terms.push(simplified); } } if terms.is_empty() { Arc::new(Expr::from(0i64)) } else if terms.len() == 1 { terms.into_iter().next().expect("non-empty") } else { let sum = Expr::add(terms); simplify(&(*sum).clone()) } } /// Compute k-th order Lie derivative L_f^k h. /// /// L_f^0 h = h /// L_f^k h = L_f(L_f^{k-1} h) #[must_use] pub fn lie_derivative_k(h: &Expr, f: &[Arc], state_vars: &[Symbol], k: usize) -> Arc { if k == 0 { return Arc::new(h.clone()); } let mut current = Arc::new(h.clone()); for _ in 0..k { current = lie_derivative(¤t, f, state_vars); } current } /// Compute all Lie derivatives L_f^0 h, L_f^1 h, …, L_f^k h. #[must_use] pub fn lie_derivatives_up_to( h: &Expr, f: &[Arc], state_vars: &[Symbol], max_order: usize, ) -> Vec> { let mut result = Vec::with_capacity(max_order + 1); let mut current = Arc::new(h.clone()); result.push(current.clone()); for _ in 0..max_order { current = lie_derivative(¤t, f, state_vars); result.push(current.clone()); } result } #[cfg(test)] mod tests { use super::*; use symclaw_core::parser::parse; fn sym(s: &str) -> Symbol { Symbol::new(s) } fn expr(s: &str) -> Arc { parse(s).expect("valid expr") } #[test] fn lie_derivative_first_order() { // h = x, f = [a*x] // L_f h = (∂h/∂x) * f = 1 * a*x = a*x let h = expr("x"); let f = vec![expr("a*x")]; let vars = vec![sym("x")]; let lh = lie_derivative(&h, &f, &vars); let s = format!("{lh}"); assert!( s.contains('a') && s.contains('x'), "L_f x should be a*x, got: {s}" ); } #[test] fn lie_derivative_chain_rule() { // h = x^2, f = [v] (ẋ = v) // L_f h = (∂x²/∂x) * v = 2x * v let h = expr("x^2"); let f = vec![expr("v")]; let vars = vec![sym("x")]; let lh = lie_derivative(&h, &f, &vars); let s = format!("{lh}"); assert!(s.contains('x'), "L_f(x²) should involve x, got: {s}"); assert!(s.contains('v'), "L_f(x²) should involve v, got: {s}"); } #[test] fn lie_derivative_constant_is_zero() { // h = 5 (constant), f = anything → L_f h = 0 let h = expr("5"); let f = vec![expr("a*x")]; let vars = vec![sym("x")]; let lh = lie_derivative(&h, &f, &vars); let s = format!("{lh}"); assert_eq!(s, "0", "L_f(5) = 0, got: {s}"); } #[test] fn lie_derivative_multivar() { // h = x + y, ẋ = a, ẏ = b // L_f h = 1*a + 1*b = a + b let h = expr("x + y"); let f = vec![expr("a"), expr("b")]; let vars = vec![sym("x"), sym("y")]; let lh = lie_derivative(&h, &f, &vars); let s = format!("{lh}"); assert!( s.contains('a') && s.contains('b'), "L_f(x+y) should be a+b, got: {s}" ); } #[test] fn lie_derivative_k_order() { // h = x, ẋ = x → L_f^k x = x (since d/dx(x)*x = x at each step) let h = expr("x"); let f = vec![expr("x")]; let vars = vec![sym("x")]; for k in 0usize..=3 { let lk = lie_derivative_k(&h, &f, &vars, k); let s = format!("{lk}"); assert!(s.contains('x'), "L_f^{k}(x) should contain x, got: {s}"); } } #[test] fn lie_derivatives_up_to_returns_correct_count() { let h = expr("x"); let f = vec![expr("a*x")]; let vars = vec![sym("x")]; let derivs = lie_derivatives_up_to(&h, &f, &vars, 4); assert_eq!( derivs.len(), 5, "up_to(4) should return 5 derivatives (0..=4)" ); } }