Initial commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user