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:
osobh
2026-08-10 07:09:36 -07:00
co-authored by Claude Sonnet 5
parent ad6405663f
commit 4aaa36a57a
305 changed files with 25537 additions and 18337 deletions
@@ -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
}
}
}
}