//! Numerical ODE solvers for Neural ODEs //! //! This module provides various numerical integration methods for solving ODEs: //! - Euler method (1st order) //! - Runge-Kutta 4th order (RK4) //! - Dormand-Prince 5th order with adaptive step size (Dopri5) //! //! Each solver implements the ODESolver trait and can be used with any ODEFunc. use super::{NeuralODEError, ODEFunc, ODEStats, Result}; use crate::prelude::*; use std::time::Instant; /// Types of ODE solvers available #[derive(Debug, Clone, PartialEq)] pub enum SolverType { /// Euler's method (1st order, simple but inaccurate) Euler, /// Runge-Kutta 4th order (good balance of accuracy and speed) RungeKutta4, /// Dormand-Prince 5th order with adaptive step size (high accuracy) Dopri5, } /// Configuration for different solver types #[derive(Debug, Clone)] pub enum SolverConfig { /// Euler method configuration Euler { /// Fixed step size step_size: f32, }, /// Runge-Kutta 4th order configuration RungeKutta4 { /// Fixed step size step_size: f32, }, /// Dormand-Prince 5th order configuration Dopri5 { /// Initial step size (will be adapted) initial_step_size: f32, /// Maximum allowed step size max_step_size: f32, /// Safety factor for step size adaptation safety_factor: f32, }, } impl Default for SolverConfig { fn default() -> Self { SolverConfig::RungeKutta4 { step_size: 0.1 } } } /// Result of a single integration step #[derive(Debug, Clone)] pub struct StepResult { /// New state after the step pub y_new: Tensor, /// Step size used pub step_size: f32, /// Estimated local error (for adaptive methods) pub error_estimate: Option, /// Whether the step was accepted pub accepted: bool, } /// Statistics collected during solving #[derive(Debug, Clone)] pub struct SolverStats { /// Number of function evaluations pub n_fe: usize, /// Number of accepted steps pub n_accepted: usize, /// Number of rejected steps pub n_rejected: usize, /// Final step size pub final_step_size: f32, /// Maximum error encountered pub max_error: f32, } impl Default for SolverStats { fn default() -> Self { Self { n_fe: 0, n_accepted: 0, n_rejected: 0, final_step_size: 0.0, max_error: 0.0, } } } /// Trait for ODE solvers pub trait ODESolver: Send + Sync { /// Solve ODE from t0 to t1 with initial condition y0 fn solve( &self, ode_func: &dyn ODEFunc, y0: &Tensor, t_span: &[f32], rtol: f32, atol: f32, ) -> Result<(Tensor, SolverStats)>; /// Take a single integration step fn step( &self, ode_func: &dyn ODEFunc, t: f32, y: &Tensor, step_size: f32, rtol: f32, atol: f32, ) -> Result; /// Get solver name for debugging fn name(&self) -> &'static str; /// Check if this is an adaptive solver fn is_adaptive(&self) -> bool; } /// Euler's method solver (1st order) pub struct EulerSolver { step_size: f32, } impl EulerSolver { /// Create a new Euler solver with fixed step size pub fn new(step_size: f32) -> Self { Self { step_size } } } impl ODESolver for EulerSolver { fn solve( &self, ode_func: &dyn ODEFunc, y0: &Tensor, t_span: &[f32], _rtol: f32, _atol: f32, ) -> Result<(Tensor, SolverStats)> { if t_span.is_empty() { return Err(NeuralODEError::InvalidInput("Empty time span".to_string())); } let device = y0.device(); let mut stats = SolverStats::default(); // Handle single time point case if t_span.len() == 1 { let result = if y0.ndim() == 1 { y0.unsqueeze(0)? // [1, state_dim] } else { y0.clone() }; return Ok((result, stats)); } let mut results = Vec::new(); let mut y_current = y0.clone(); let mut t_current = t_span[0]; // Add initial condition results.push(y_current.clone()); for &t_target in &t_span[1..] { while t_current < t_target { let remaining_time = t_target - t_current; let actual_step = self.step_size.min(remaining_time); let step_result = self.step(ode_func, t_current, &y_current, actual_step, 0.0, 0.0)?; y_current = step_result.y_new; t_current += actual_step; stats.n_fe += 1; stats.n_accepted += 1; } results.push(y_current.clone()); } stats.final_step_size = self.step_size; // Stack results along time dimension let output = Tensor::stack(&results, 0)?; Ok((output, stats)) } fn step( &self, ode_func: &dyn ODEFunc, t: f32, y: &Tensor, step_size: f32, _rtol: f32, _atol: f32, ) -> Result { // Euler step: y_{n+1} = y_n + h * f(t_n, y_n) let dy_dt = ode_func.forward(t, y)?; let y_new = y.add(&dy_dt.mul_scalar(step_size)?)?; Ok(StepResult { y_new, step_size, error_estimate: None, // Euler doesn't provide error estimates accepted: true, }) } fn name(&self) -> &'static str { "Euler" } fn is_adaptive(&self) -> bool { false } } /// Runge-Kutta 4th order solver pub struct RungeKutta4Solver { step_size: f32, } impl RungeKutta4Solver { /// Create a new RK4 solver with fixed step size pub fn new(step_size: f32) -> Self { Self { step_size } } } impl ODESolver for RungeKutta4Solver { fn solve( &self, ode_func: &dyn ODEFunc, y0: &Tensor, t_span: &[f32], _rtol: f32, _atol: f32, ) -> Result<(Tensor, SolverStats)> { if t_span.is_empty() { return Err(NeuralODEError::InvalidInput("Empty time span".to_string())); } let device = y0.device(); let mut stats = SolverStats::default(); // Handle single time point case if t_span.len() == 1 { let result = if y0.ndim() == 1 { y0.unsqueeze(0)? // [1, state_dim] } else { y0.clone() }; return Ok((result, stats)); } let mut results = Vec::new(); let mut y_current = y0.clone(); let mut t_current = t_span[0]; // Add initial condition results.push(y_current.clone()); for &t_target in &t_span[1..] { while t_current < t_target { let remaining_time = t_target - t_current; let actual_step = self.step_size.min(remaining_time); let step_result = self.step(ode_func, t_current, &y_current, actual_step, 0.0, 0.0)?; y_current = step_result.y_new; t_current += actual_step; stats.n_fe += 4; // RK4 uses 4 function evaluations per step stats.n_accepted += 1; } results.push(y_current.clone()); } stats.final_step_size = self.step_size; // Stack results along time dimension let output = Tensor::stack(&results, 0)?; Ok((output, stats)) } fn step( &self, ode_func: &dyn ODEFunc, t: f32, y: &Tensor, step_size: f32, _rtol: f32, _atol: f32, ) -> Result { let h = step_size; let h_half = h * 0.5; // RK4 method: // k1 = f(t, y) // k2 = f(t + h/2, y + h/2 * k1) // k3 = f(t + h/2, y + h/2 * k2) // k4 = f(t + h, y + h * k3) // y_{n+1} = y_n + h/6 * (k1 + 2*k2 + 2*k3 + k4) let k1 = ode_func.forward(t, y)?; let y_k2 = y.add(&k1.mul_scalar(h_half)?)?; let k2 = ode_func.forward(t + h_half, &y_k2)?; let y_k3 = y.add(&k2.mul_scalar(h_half)?)?; let k3 = ode_func.forward(t + h_half, &y_k3)?; let y_k4 = y.add(&k3.mul_scalar(h)?)?; let k4 = ode_func.forward(t + h, &y_k4)?; // Combine: y_new = y + h/6 * (k1 + 2*k2 + 2*k3 + k4) let k_combined = k1 .add(&k2.mul_scalar(2.0)?)? .add(&k3.mul_scalar(2.0)?)? .add(&k4)?; let y_new = y.add(&k_combined.mul_scalar(h / 6.0)?)?; Ok(StepResult { y_new, step_size, error_estimate: None, // Standard RK4 doesn't provide error estimates accepted: true, }) } fn name(&self) -> &'static str { "RungeKutta4" } fn is_adaptive(&self) -> bool { false } } /// Dormand-Prince 5th order adaptive solver pub struct Dopri5Solver { initial_step_size: f32, max_step_size: f32, safety_factor: f32, } impl Dopri5Solver { /// Create a new Dopri5 solver with adaptive step size pub fn new(initial_step_size: f32, max_step_size: f32, safety_factor: f32) -> Self { Self { initial_step_size, max_step_size, safety_factor, } } /// Compute error estimate for step size control fn error_estimate(&self, k_vals: &[Tensor]) -> Result { // Error estimate using embedded 4th order method // This is a simplified version - in practice would use Dormand-Prince coefficients let error_coeffs = [ 1.0 / 360.0, 0.0, -128.0 / 4275.0, -2197.0 / 75240.0, 1.0 / 50.0, 2.0 / 55.0, ]; let mut error_vec = k_vals[0].mul_scalar(error_coeffs[0])?; for (i, k) in k_vals.iter().enumerate().skip(1) { if error_coeffs[i] != 0.0 { error_vec = error_vec.add(&k.mul_scalar(error_coeffs[i])?)?; } } let error_norm = error_vec .pow_scalar(2.0)? .sum(None)? .sqrt()? .to_scalar::()?; Ok(error_norm) } } impl ODESolver for Dopri5Solver { fn solve( &self, ode_func: &dyn ODEFunc, y0: &Tensor, t_span: &[f32], rtol: f32, atol: f32, ) -> Result<(Tensor, SolverStats)> { if t_span.is_empty() { return Err(NeuralODEError::InvalidInput("Empty time span".to_string())); } let mut stats = SolverStats::default(); // Handle single time point case if t_span.len() == 1 { let result = if y0.ndim() == 1 { y0.unsqueeze(0)? // [1, state_dim] } else { y0.clone() }; return Ok((result, stats)); } let mut results = Vec::new(); let mut y_current = y0.clone(); let mut t_current = t_span[0]; let mut current_step_size = self.initial_step_size; // Add initial condition results.push(y_current.clone()); for &t_target in &t_span[1..] { while t_current < t_target { let remaining_time = t_target - t_current; let proposed_step = current_step_size .min(self.max_step_size) .min(remaining_time); let step_result = self.step(ode_func, t_current, &y_current, proposed_step, rtol, atol)?; stats.n_fe += 6; // Dopri5 uses 6 function evaluations per step attempt if step_result.accepted { y_current = step_result.y_new; t_current += step_result.step_size; stats.n_accepted += 1; // Update step size based on error estimate if let Some(error) = step_result.error_estimate { let tolerance = rtol * y_current.abs()?.max()?.to_scalar::()? + atol; let ratio = (tolerance / error.max(1e-12)).powf(0.2); // 5th order -> 1/5 = 0.2 current_step_size = (step_result.step_size * self.safety_factor * ratio) .min(self.max_step_size) .max(1e-8); stats.max_error = stats.max_error.max(error); } } else { stats.n_rejected += 1; // Reduce step size for rejected step current_step_size *= 0.5; if current_step_size < 1e-8 { return Err(NeuralODEError::Convergence( "Step size became too small".to_string(), )); } } } results.push(y_current.clone()); } stats.final_step_size = current_step_size; // Stack results along time dimension let output = Tensor::stack(&results, 0)?; Ok((output, stats)) } fn step( &self, ode_func: &dyn ODEFunc, t: f32, y: &Tensor, step_size: f32, rtol: f32, atol: f32, ) -> Result { let h = step_size; // Dormand-Prince 5(4) method - simplified version // In practice, would use the full Butcher tableau let k1 = ode_func.forward(t, y)?; let y2 = y.add(&k1.mul_scalar(h * 0.2)?)?; let k2 = ode_func.forward(t + h * 0.2, &y2)?; let y3 = y .add(&k1.mul_scalar(h * 3.0 / 40.0)?)? .add(&k2.mul_scalar(h * 9.0 / 40.0)?)?; let k3 = ode_func.forward(t + h * 0.3, &y3)?; let y4 = y .add(&k1.mul_scalar(h * 44.0 / 45.0)?)? .sub(&k2.mul_scalar(h * 56.0 / 15.0)?)? .add(&k3.mul_scalar(h * 32.0 / 9.0)?)?; let k4 = ode_func.forward(t + h * 0.8, &y4)?; let y5 = y .add(&k1.mul_scalar(h * 19372.0 / 6561.0)?)? .sub(&k2.mul_scalar(h * 25360.0 / 2187.0)?)? .add(&k3.mul_scalar(h * 64448.0 / 6561.0)?)? .sub(&k4.mul_scalar(h * 212.0 / 729.0)?)?; let k5 = ode_func.forward(t + h * 8.0 / 9.0, &y5)?; let y6 = y .add(&k1.mul_scalar(h * 9017.0 / 3168.0)?)? .sub(&k2.mul_scalar(h * 355.0 / 33.0)?)? .add(&k3.mul_scalar(h * 46732.0 / 5247.0)?)? .add(&k4.mul_scalar(h * 49.0 / 176.0)?)? .sub(&k5.mul_scalar(h * 5103.0 / 18656.0)?)?; let k6 = ode_func.forward(t + h, &y6)?; // 5th order solution let y_new = y.add( &k1.mul_scalar(h * 35.0 / 384.0)? .add(&k3.mul_scalar(h * 500.0 / 1113.0)?)? .add(&k4.mul_scalar(h * 125.0 / 192.0)?)? .sub(&k5.mul_scalar(h * 2187.0 / 6784.0)?)? .add(&k6.mul_scalar(h * 11.0 / 84.0)?)?, )?; // Error estimate (4th order vs 5th order) let k_vals = vec![k1, k2, k3, k4, k5, k6]; let error_estimate = self.error_estimate(&k_vals)?; // Check if step should be accepted let tolerance = rtol * y.abs()?.max()?.to_scalar::()? + atol; let accepted = error_estimate <= tolerance; Ok(StepResult { y_new, step_size, error_estimate: Some(error_estimate), accepted, }) } fn name(&self) -> &'static str { "Dopri5" } fn is_adaptive(&self) -> bool { true } } /// Factory function to create solvers from configuration pub fn create_solver(solver_type: SolverType, config: SolverConfig) -> Result> { match (solver_type, config) { (SolverType::Euler, SolverConfig::Euler { step_size }) => { Ok(Box::new(EulerSolver::new(step_size))) } (SolverType::RungeKutta4, SolverConfig::RungeKutta4 { step_size }) => { Ok(Box::new(RungeKutta4Solver::new(step_size))) } ( SolverType::Dopri5, SolverConfig::Dopri5 { initial_step_size, max_step_size, safety_factor, }, ) => Ok(Box::new(Dopri5Solver::new( initial_step_size, max_step_size, safety_factor, ))), _ => Err(NeuralODEError::Config( "Solver type and config mismatch".to_string(), )), } } #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::*; use crate::neural_ode::ode_func::LinearDynamics; #[test] fn test_euler_solver() { let device = Device::cpu(); let dynamics = LinearDynamics::decay(1.0, 1, device.clone()).unwrap(); let solver = EulerSolver::new(0.1); let y0 = Tensor::ones([1], &device).unwrap(); let t_span = vec![0.0, 0.5]; let (result, stats) = solver.solve(&dynamics, &y0, &t_span, 1e-3, 1e-6).unwrap(); assert_eq!(result.shape(), &[2, 1]); // [time_steps, state_dim] assert!(stats.n_fe > 0); assert_eq!(stats.n_rejected, 0); // Fixed step solver let result_slice = result.to_vec().unwrap(); assert!((result_slice[0] - 1.0).abs() < 1e-6); // Initial condition assert!(result_slice[1] < 1.0); // Should decay } #[test] fn test_rk4_solver() { let device = Device::cpu(); let dynamics = LinearDynamics::decay(1.0, 1, device.clone()).unwrap(); let solver = RungeKutta4Solver::new(0.1); let y0 = Tensor::ones([1], &device).unwrap(); let t_span = vec![0.0, 1.0]; let (result, stats) = solver.solve(&dynamics, &y0, &t_span, 1e-3, 1e-6).unwrap(); assert_eq!(result.shape(), &[2, 1]); let result_slice = result.to_vec().unwrap(); let expected = (-1.0_f32).exp(); // Analytical solution: e^(-t) // RK4 should be more accurate than Euler assert!((result_slice[1] - expected).abs() < 0.01); } #[test] fn test_dopri5_solver() { let device = Device::cpu(); let dynamics = LinearDynamics::decay(1.0, 1, device.clone()).unwrap(); let solver = Dopri5Solver::new(0.01, 0.1, 0.9); let y0 = Tensor::ones([1], &device).unwrap(); let t_span = vec![0.0, 1.0]; let (result, stats) = solver.solve(&dynamics, &y0, &t_span, 1e-5, 1e-8).unwrap(); assert_eq!(result.shape(), &[2, 1]); let result_slice = result.to_vec().unwrap(); let expected = (-1.0_f32).exp(); // Adaptive solver should be very accurate assert!((result_slice[1] - expected).abs() < 1e-4); // Should have some function evaluations assert!(stats.n_fe > 0); } #[test] fn test_solver_factory() { let euler = create_solver(SolverType::Euler, SolverConfig::Euler { step_size: 0.1 }).unwrap(); assert_eq!(euler.name(), "Euler"); assert!(!euler.is_adaptive()); let rk4 = create_solver( SolverType::RungeKutta4, SolverConfig::RungeKutta4 { step_size: 0.05 }, ) .unwrap(); assert_eq!(rk4.name(), "RungeKutta4"); assert!(!rk4.is_adaptive()); let dopri5 = create_solver( SolverType::Dopri5, SolverConfig::Dopri5 { initial_step_size: 0.01, max_step_size: 0.1, safety_factor: 0.9, }, ) .unwrap(); assert_eq!(dopri5.name(), "Dopri5"); assert!(dopri5.is_adaptive()); } #[test] fn test_single_time_point() { let device = Device::cpu(); let dynamics = LinearDynamics::decay(1.0, 2, device.clone()).unwrap(); let solver = EulerSolver::new(0.1); let y0 = Tensor::ones([2], &device).unwrap(); let t_span = vec![0.0]; let (result, _) = solver.solve(&dynamics, &y0, &t_span, 1e-3, 1e-6).unwrap(); assert_eq!(result.shape(), &[1, 2]); // Single time point let result_slice = result.to_vec().unwrap(); assert_eq!(result_slice, vec![1.0, 1.0]); // Should return initial condition } #[test] fn test_batch_solving() { let device = Device::cpu(); let dynamics = LinearDynamics::decay(1.0, 2, device.clone()).unwrap(); let solver = RungeKutta4Solver::new(0.1); // Batch of initial conditions let y0_batch = Tensor::ones([3, 2], &device).unwrap(); let t_span = vec![0.0, 0.5]; // For now, we need to solve each sample individually // In a full implementation, we'd support batch solving natively for i in 0..3 { let y0_single = y0_batch.narrow(0, i, 1).unwrap().squeeze(Some(0)).unwrap(); let (result, _) = solver .solve(&dynamics, &y0_single, &t_span, 1e-3, 1e-6) .unwrap(); assert_eq!(result.shape(), &[2, 2]); let result_slice = result.to_vec().unwrap(); assert!(result_slice[0] >= result_slice[2]); // Should decay over time } } }