Whole-workspace rustfmt pass picked up while iterating on Mamba GPU backward work. Verified formatting-only via diff sampling; no logic changed. Co-Authored-By: Claude Sonnet 5 <[email protected]>
721 lines
21 KiB
Rust
721 lines
21 KiB
Rust
//! 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<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(
|
|
&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<StepResult>;
|
|
|
|
/// 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<StepResult> {
|
|
// 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<StepResult> {
|
|
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(None)?
|
|
.sqrt()?
|
|
.to_scalar::<f32>()?;
|
|
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::<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(
|
|
&self,
|
|
ode_func: &dyn ODEFunc,
|
|
t: f32,
|
|
y: &Tensor,
|
|
step_size: f32,
|
|
rtol: f32,
|
|
atol: f32,
|
|
) -> Result<StepResult> {
|
|
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().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
|
|
}
|
|
}
|
|
}
|