Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,712 @@
//! 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 crate::prelude::*;
use super::{ODEFunc, Result, NeuralODEError, ODEStats};
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<f32>,
/// 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<F>(
&self,
ode_func: &F,
y0: &Tensor,
t_span: &[f32],
rtol: f32,
atol: f32,
) -> Result<(Tensor, SolverStats)>
where
F: ODEFunc;
/// Take a single integration step
fn step<F>(
&self,
ode_func: &F,
t: f32,
y: &Tensor,
step_size: f32,
rtol: f32,
atol: f32,
) -> Result<StepResult>
where
F: ODEFunc;
/// 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<F>(
&self,
ode_func: &F,
y0: &Tensor,
t_span: &[f32],
_rtol: f32,
_atol: f32,
) -> Result<(Tensor, SolverStats)>
where
F: ODEFunc,
{
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<F>(
&self,
ode_func: &F,
t: f32,
y: &Tensor,
step_size: f32,
_rtol: f32,
_atol: f32,
) -> Result<StepResult>
where
F: ODEFunc,
{
// 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<F>(
&self,
ode_func: &F,
y0: &Tensor,
t_span: &[f32],
_rtol: f32,
_atol: f32,
) -> Result<(Tensor, SolverStats)>
where
F: ODEFunc,
{
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<F>(
&self,
ode_func: &F,
t: f32,
y: &Tensor,
step_size: f32,
_rtol: f32,
_atol: f32,
) -> Result<StepResult>
where
F: ODEFunc,
{
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<f32> {
// 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()?.sqrt()?.to_scalar::<f32>()?;
Ok(error_norm)
}
}
impl ODESolver for Dopri5Solver {
fn solve<F>(
&self,
ode_func: &F,
y0: &Tensor,
t_span: &[f32],
rtol: f32,
atol: f32,
) -> Result<(Tensor, SolverStats)>
where
F: ODEFunc,
{
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::<f32>()? + 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<F>(
&self,
ode_func: &F,
t: f32,
y: &Tensor,
step_size: f32,
rtol: f32,
atol: f32,
) -> Result<StepResult>
where
F: ODEFunc,
{
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::<f32>()? + 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<Box<dyn ODESolver>> {
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::<f32>().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::<f32>().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::<f32>().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::<f32>().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.slice(0, i..i+1).unwrap().squeeze(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::<f32>().unwrap();
assert!(result_slice[0] >= result_slice[2]); // Should decay over time
}
}
}