style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
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]>
This commit is contained in:
@@ -7,8 +7,8 @@
|
||||
//!
|
||||
//! Each solver implements the ODESolver trait and can be used with any ODEFunc.
|
||||
|
||||
use super::{NeuralODEError, ODEFunc, ODEStats, Result};
|
||||
use crate::prelude::*;
|
||||
use super::{ODEFunc, Result, NeuralODEError, ODEStats};
|
||||
use std::time::Instant;
|
||||
|
||||
/// Types of ODE solvers available
|
||||
@@ -142,15 +142,14 @@ impl ODESolver for EulerSolver {
|
||||
t_span: &[f32],
|
||||
_rtol: f32,
|
||||
_atol: f32,
|
||||
) -> Result<(Tensor, SolverStats)>
|
||||
{
|
||||
) -> 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 {
|
||||
@@ -173,20 +172,21 @@ impl ODESolver for EulerSolver {
|
||||
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)?;
|
||||
|
||||
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))
|
||||
@@ -200,8 +200,7 @@ impl ODESolver for EulerSolver {
|
||||
step_size: f32,
|
||||
_rtol: f32,
|
||||
_atol: f32,
|
||||
) -> Result<StepResult>
|
||||
{
|
||||
) -> 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)?)?;
|
||||
@@ -243,15 +242,14 @@ impl ODESolver for RungeKutta4Solver {
|
||||
t_span: &[f32],
|
||||
_rtol: f32,
|
||||
_atol: f32,
|
||||
) -> Result<(Tensor, SolverStats)>
|
||||
{
|
||||
) -> 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 {
|
||||
@@ -274,20 +272,21 @@ impl ODESolver for RungeKutta4Solver {
|
||||
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)?;
|
||||
|
||||
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))
|
||||
@@ -301,8 +300,7 @@ impl ODESolver for RungeKutta4Solver {
|
||||
step_size: f32,
|
||||
_rtol: f32,
|
||||
_atol: f32,
|
||||
) -> Result<StepResult>
|
||||
{
|
||||
) -> Result<StepResult> {
|
||||
let h = step_size;
|
||||
let h_half = h * 0.5;
|
||||
|
||||
@@ -314,22 +312,22 @@ impl ODESolver for RungeKutta4Solver {
|
||||
// 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 {
|
||||
@@ -370,16 +368,27 @@ impl Dopri5Solver {
|
||||
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 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>()?;
|
||||
|
||||
let error_norm = error_vec
|
||||
.pow_scalar(2.0)?
|
||||
.sum(None)?
|
||||
.sqrt()?
|
||||
.to_scalar::<f32>()?;
|
||||
Ok(error_norm)
|
||||
}
|
||||
}
|
||||
@@ -392,14 +401,13 @@ impl ODESolver for Dopri5Solver {
|
||||
t_span: &[f32],
|
||||
rtol: f32,
|
||||
atol: f32,
|
||||
) -> Result<(Tensor, SolverStats)>
|
||||
{
|
||||
) -> 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 {
|
||||
@@ -421,17 +429,20 @@ impl ODESolver for Dopri5Solver {
|
||||
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 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)?;
|
||||
|
||||
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;
|
||||
@@ -439,27 +450,27 @@ impl ODESolver for Dopri5Solver {
|
||||
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()
|
||||
"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))
|
||||
@@ -473,52 +484,56 @@ impl ODESolver for Dopri5Solver {
|
||||
step_size: f32,
|
||||
rtol: f32,
|
||||
atol: f32,
|
||||
) -> Result<StepResult>
|
||||
{
|
||||
) -> 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 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 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 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)?)?
|
||||
&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;
|
||||
@@ -549,14 +564,21 @@ pub fn create_solver(solver_type: SolverType, config: SolverConfig) -> Result<Bo
|
||||
(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)))
|
||||
}
|
||||
(
|
||||
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()
|
||||
))
|
||||
"Solver type and config mismatch".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -570,16 +592,16 @@ mod tests {
|
||||
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
|
||||
@@ -590,17 +612,17 @@ mod tests {
|
||||
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);
|
||||
}
|
||||
@@ -610,48 +632,48 @@ mod tests {
|
||||
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();
|
||||
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();
|
||||
SolverConfig::RungeKutta4 { step_size: 0.05 },
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(rk4.name(), "RungeKutta4");
|
||||
assert!(!rk4.is_adaptive());
|
||||
|
||||
let dopri5 = create_solver(
|
||||
SolverType::Dopri5,
|
||||
SolverConfig::Dopri5 {
|
||||
SolverConfig::Dopri5 {
|
||||
initial_step_size: 0.01,
|
||||
max_step_size: 0.1,
|
||||
safety_factor: 0.9
|
||||
}
|
||||
).unwrap();
|
||||
safety_factor: 0.9,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(dopri5.name(), "Dopri5");
|
||||
assert!(dopri5.is_adaptive());
|
||||
}
|
||||
@@ -661,12 +683,12 @@ mod tests {
|
||||
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
|
||||
@@ -677,20 +699,22 @@ mod tests {
|
||||
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();
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user