Files
rustytorch/crates/training/rtx-transformers/tests/neural_ode_tests.rs
T
2026-03-04 00:08:42 +00:00

980 lines
33 KiB
Rust

//! Comprehensive test suite for Neural ODE implementation following TDD methodology
//!
//! NOTE: Disabled until Neural ODE API is fully implemented
//!
//! This test suite covers all aspects of Neural ODEs:
//! - Basic ODE solving with different numerical methods
//! - Gradient computation via adjoint sensitivity method
//! - Continuous normalizing flows (CNF) for density estimation
//! - Time series modeling with irregular sampling
//! - Performance comparisons between solvers
//! - Augmented Neural ODEs with additional dimensions
#![cfg(feature = "disabled_tests")]
use approx::assert_relative_eq;
use rtx_transformers::neural_ode::{
AdjointConfig, AdjointMethod, AugmentedNeuralODE, CNFConfig, ContinuousNormalizingFlow,
NeuralODE, NeuralODEConfig, ODEFunc, ODESolver, SolverConfig, SolverType,
};
use rtx_transformers::prelude::*;
use std::collections::HashMap;
/// Test fixture for simple 2D dynamics: dx/dt = -x
struct SimpleLinearDynamics {
device: Device,
}
impl SimpleLinearDynamics {
fn new(device: Device) -> Self {
Self { device }
}
}
impl ODEFunc for SimpleLinearDynamics {
fn forward(&self, t: f32, y: &Tensor) -> Result<Tensor> {
// dx/dt = -x (simple exponential decay)
y.neg()
}
fn parameters(&self) -> Vec<Tensor> {
vec![] // No learnable parameters for this simple case
}
}
/// Test fixture for neural network dynamics
struct NeuralNetworkDynamics {
linear1: Tensor, // Weight matrix
bias1: Tensor, // Bias vector
linear2: Tensor, // Second layer
bias2: Tensor,
device: Device,
}
impl NeuralNetworkDynamics {
fn new(input_dim: usize, hidden_dim: usize, device: Device) -> Result<Self> {
let linear1 = Tensor::randn([hidden_dim, input_dim], &device)?;
let bias1 = Tensor::zeros([hidden_dim], &device)?;
let linear2 = Tensor::randn([input_dim, hidden_dim], &device)?;
let bias2 = Tensor::zeros([input_dim], &device)?;
Ok(Self {
linear1,
bias1,
linear2,
bias2,
device,
})
}
}
impl ODEFunc for NeuralNetworkDynamics {
fn forward(&self, t: f32, y: &Tensor) -> Result<Tensor> {
// Two-layer neural network: tanh(W1*y + b1) -> W2*h + b2
let h1 = y.matmul(&self.linear1.t()?)?.add(&self.bias1)?;
let h1_activated = h1.tanh()?;
let output = h1_activated.matmul(&self.linear2.t()?)?.add(&self.bias2)?;
Ok(output)
}
fn parameters(&self) -> Vec<Tensor> {
vec![
self.linear1.clone(),
self.bias1.clone(),
self.linear2.clone(),
self.bias2.clone(),
]
}
}
#[cfg(test)]
mod basic_ode_solving {
use super::*;
#[test]
fn test_euler_solver_simple_dynamics() {
let device = Device::cpu();
let dynamics = SimpleLinearDynamics::new(device.clone());
let config = NeuralODEConfig {
solver: SolverType::Euler,
solver_config: SolverConfig::Euler { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-3,
atol: 1e-6,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
// Initial condition: y0 = [1.0, 2.0]
let y0 = Tensor::from_slice(&[1.0, 2.0], [2], &device).unwrap();
let t_span = vec![0.0, 1.0];
// Should solve dx/dt = -x, so x(t) = x0 * exp(-t)
let result = neural_ode.forward(&y0, &t_span).unwrap();
// Expected: [exp(-1), 2*exp(-1)] ≈ [0.368, 0.736]
let expected_val = (-1.0_f32).exp();
let result_slice = result.to_vec::<f32>().unwrap();
assert_relative_eq!(result_slice[0], expected_val, epsilon = 0.1);
assert_relative_eq!(result_slice[1], 2.0 * expected_val, epsilon = 0.1);
}
#[test]
fn test_rk4_solver_higher_accuracy() {
let device = Device::cpu();
let dynamics = SimpleLinearDynamics::new(device.clone());
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-6,
atol: 1e-9,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = Tensor::from_slice(&[1.0, 2.0], [2], &device).unwrap();
let t_span = vec![0.0, 1.0];
let result = neural_ode.forward(&y0, &t_span).unwrap();
let expected_val = (-1.0_f32).exp();
let result_slice = result.to_vec::<f32>().unwrap();
// RK4 should be much more accurate than Euler
assert_relative_eq!(result_slice[0], expected_val, epsilon = 1e-3);
assert_relative_eq!(result_slice[1], 2.0 * expected_val, epsilon = 1e-3);
}
#[test]
fn test_adaptive_solver_error_control() {
let device = Device::cpu();
let dynamics = SimpleLinearDynamics::new(device.clone());
let config = NeuralODEConfig {
solver: SolverType::Dopri5,
solver_config: SolverConfig::Dopri5 {
initial_step_size: 0.01,
max_step_size: 0.5,
safety_factor: 0.9,
},
adjoint_config: None,
rtol: 1e-6,
atol: 1e-9,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = Tensor::from_slice(&[1.0, 2.0], [2], &device).unwrap();
let t_span = vec![0.0, 1.0];
let result = neural_ode.forward(&y0, &t_span).unwrap();
let expected_val = (-1.0_f32).exp();
let result_slice = result.to_vec::<f32>().unwrap();
// Adaptive solver should achieve very high accuracy
assert_relative_eq!(result_slice[0], expected_val, epsilon = 1e-5);
assert_relative_eq!(result_slice[1], 2.0 * expected_val, epsilon = 1e-5);
}
#[test]
fn test_neural_network_dynamics() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(2, 4, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = Tensor::randn([2], &device).unwrap();
let t_span = vec![0.0, 0.5, 1.0];
let result = neural_ode.forward(&y0, &t_span).unwrap();
// Check that we get reasonable output shape and values
assert_eq!(result.shape(), &[3, 2]); // [time_steps, state_dim]
assert!(!result.to_vec::<f32>().unwrap().iter().any(|&x| x.is_nan()));
}
}
#[cfg(test)]
mod gradient_computation {
use super::*;
#[test]
fn test_adjoint_method_gradient_accuracy() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(2, 4, device.clone()).unwrap();
let adjoint_config = AdjointConfig {
solver_type: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
rtol: 1e-5,
atol: 1e-8,
};
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
adjoint_config: Some(adjoint_config),
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = tensor_with_grad(Tensor::randn([2], &device).unwrap());
let t_span = vec![0.0, 1.0];
let result = neural_ode.forward(&y0, &t_span).unwrap();
let loss = result.pow_scalar(2.0)?.sum()?;
// Compute gradients using adjoint method
let gradients = backward(loss.node_id(), None).unwrap();
// Check that gradients exist for all parameters
let params = neural_ode.parameters();
for param in params {
assert!(gradients.contains_key(&param.node_id()));
let grad = &gradients[&param.node_id()];
assert_eq!(grad.shape(), param.shape());
}
// Gradients should not be all zeros (assuming non-trivial dynamics)
let total_grad_norm: f32 = gradients
.values()
.map(|g| {
g.pow_scalar(2.0)
.unwrap()
.sum()
.unwrap()
.to_scalar::<f32>()
.unwrap()
})
.sum();
assert!(total_grad_norm > 1e-6);
}
#[test]
fn test_gradient_checkpointing_memory_efficiency() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(10, 20, device.clone()).unwrap();
let adjoint_config = AdjointConfig {
solver_type: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
rtol: 1e-4,
atol: 1e-7,
};
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: Some(adjoint_config),
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = tensor_with_grad(Tensor::randn([10], &device).unwrap());
let t_span = vec![0.0, 0.5, 1.0, 1.5, 2.0]; // Many time steps
let result = neural_ode.forward(&y0, &t_span).unwrap();
let loss = result.pow_scalar(2.0)?.sum()?;
// This should complete without running out of memory
let gradients = backward(loss.node_id(), None).unwrap();
assert!(!gradients.is_empty());
}
#[test]
fn test_numerical_gradient_verification() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(2, 3, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: Some(AdjointConfig {
solver_type: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
rtol: 1e-5,
atol: 1e-8,
}),
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
// Use gradient checker to verify adjoint accuracy
let checker = GradientChecker::new(1e-5, 1e-3);
let y0 = tensor_with_grad(Tensor::randn([2], &device).unwrap());
let t_span = vec![0.0, 0.5];
let check_result = checker.check_gradients(&neural_ode, &y0, &t_span).unwrap();
// All gradient checks should pass
assert!(check_result.all_passed);
assert!(check_result.max_error < 1e-3);
}
}
#[cfg(test)]
mod continuous_normalizing_flows {
use super::*;
struct CNFDynamics {
linear: Tensor,
bias: Tensor,
device: Device,
}
impl CNFDynamics {
fn new(dim: usize, device: Device) -> Result<Self> {
let linear = Tensor::randn([dim, dim], &device)?;
let bias = Tensor::zeros([dim], &device)?;
Ok(Self {
linear,
bias,
device,
})
}
}
impl ODEFunc for CNFDynamics {
fn forward(&self, t: f32, y: &Tensor) -> Result<Tensor> {
let h = y.matmul(&self.linear.t()?)?.add(&self.bias)?;
h.tanh() // Bounded dynamics for stability
}
fn parameters(&self) -> Vec<Tensor> {
vec![self.linear.clone(), self.bias.clone()]
}
}
#[test]
fn test_cnf_density_estimation() {
let device = Device::cpu();
let dynamics = CNFDynamics::new(2, device.clone()).unwrap();
let cnf_config = CNFConfig {
solver: SolverType::Dopri5,
solver_config: SolverConfig::Dopri5 {
initial_step_size: 0.01,
max_step_size: 0.1,
safety_factor: 0.9,
},
compute_divergence: true,
hutchinson_trace_estimator: true,
rtol: 1e-5,
atol: 1e-8,
};
let cnf = ContinuousNormalizingFlow::new(Box::new(dynamics), cnf_config, &device).unwrap();
// Sample from base distribution (standard Gaussian)
let z0 = Tensor::randn([100, 2], &device).unwrap();
let t_span = vec![0.0, 1.0];
let (x1, log_det_jac) = cnf.forward(&z0, &t_span).unwrap();
// Check output shapes
assert_eq!(x1.shape(), &[100, 2]);
assert_eq!(log_det_jac.shape(), &[100]);
// Log determinant should be finite
let log_det_slice = log_det_jac.to_vec::<f32>().unwrap();
assert!(log_det_slice.iter().all(|&x| x.is_finite()));
}
#[test]
fn test_cnf_invertibility() {
let device = Device::cpu();
let dynamics = CNFDynamics::new(2, device.clone()).unwrap();
let cnf_config = CNFConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
compute_divergence: false,
hutchinson_trace_estimator: false,
rtol: 1e-4,
atol: 1e-7,
};
let cnf = ContinuousNormalizingFlow::new(Box::new(dynamics), cnf_config, &device).unwrap();
let x0 = Tensor::randn([10, 2], &device).unwrap();
let t_span = vec![0.0, 1.0];
// Forward transformation
let (x1, _) = cnf.forward(&x0, &t_span).unwrap();
// Inverse transformation (reverse time)
let t_span_reverse = vec![1.0, 0.0];
let (x0_reconstructed, _) = cnf.forward(&x1, &t_span_reverse).unwrap();
// Should reconstruct original input (within numerical tolerance)
let x0_slice = x0.to_vec::<f32>().unwrap();
let x0_recon_slice = x0_reconstructed.to_vec::<f32>().unwrap();
for (orig, recon) in x0_slice.iter().zip(x0_recon_slice.iter()) {
assert_relative_eq!(*orig, *recon, epsilon = 1e-2);
}
}
#[test]
fn test_cnf_training_loss() {
let device = Device::cpu();
let dynamics = CNFDynamics::new(2, device.clone()).unwrap();
let cnf_config = CNFConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
compute_divergence: true,
hutchinson_trace_estimator: true,
rtol: 1e-4,
atol: 1e-7,
};
let cnf = ContinuousNormalizingFlow::new(Box::new(dynamics), cnf_config, &device).unwrap();
// Target data (e.g., samples from some distribution)
let target_data = Tensor::randn([50, 2], &device).unwrap();
// Base distribution samples
let z0 = Tensor::randn([50, 2], &device).unwrap();
let t_span = vec![0.0, 1.0];
let (x1, log_det_jac) = cnf.forward(&z0, &t_span).unwrap();
// Compute negative log-likelihood loss
let reconstruction_loss = target_data.sub(&x1)?.pow_scalar(2.0)?.mean()?;
let jacobian_loss = log_det_jac.mean()?;
let total_loss = reconstruction_loss.sub(&jacobian_loss)?;
// Loss should be computable and finite
let loss_val = total_loss.to_scalar::<f32>().unwrap();
assert!(loss_val.is_finite());
}
}
#[cfg(test)]
mod time_series_modeling {
use super::*;
struct TimeSeriesDynamics {
gru_weight: Tensor,
gru_bias: Tensor,
output_weight: Tensor,
hidden_state: Option<Tensor>,
device: Device,
}
impl TimeSeriesDynamics {
fn new(input_dim: usize, hidden_dim: usize, device: Device) -> Result<Self> {
let gru_weight = Tensor::randn([3 * hidden_dim, input_dim + hidden_dim], &device)?;
let gru_bias = Tensor::zeros([3 * hidden_dim], &device)?;
let output_weight = Tensor::randn([input_dim, hidden_dim], &device)?;
Ok(Self {
gru_weight,
gru_bias,
output_weight,
hidden_state: None,
device,
})
}
}
impl ODEFunc for TimeSeriesDynamics {
fn forward(&self, t: f32, y: &Tensor) -> Result<Tensor> {
// Simple time-dependent dynamics: dy/dt = -y + sin(t)
let time_component = Tensor::full_like(y, t.sin(), &self.device)?;
let decay_component = y.neg()?;
decay_component.add(&time_component)
}
fn parameters(&self) -> Vec<Tensor> {
vec![
self.gru_weight.clone(),
self.gru_bias.clone(),
self.output_weight.clone(),
]
}
}
#[test]
fn test_irregular_time_series() {
let device = Device::cpu();
let dynamics = TimeSeriesDynamics::new(1, 4, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::Dopri5,
solver_config: SolverConfig::Dopri5 {
initial_step_size: 0.01,
max_step_size: 0.1,
safety_factor: 0.9,
},
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = Tensor::from_slice(&[1.0], [1], &device).unwrap();
// Irregular time points
let t_span = vec![0.0, 0.1, 0.3, 0.8, 1.5, 2.1, 3.0];
let result = neural_ode.forward(&y0, &t_span).unwrap();
// Check output shape matches time points
assert_eq!(result.shape(), &[7, 1]); // [time_steps, state_dim]
// Values should be finite and reasonable
let result_slice = result.to_vec::<f32>().unwrap();
assert!(result_slice.iter().all(|&x| x.is_finite()));
}
#[test]
fn test_time_dependent_dynamics() {
let device = Device::cpu();
let dynamics = TimeSeriesDynamics::new(2, 4, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = Tensor::zeros([2], &device).unwrap();
let t_span = vec![0.0, std::f32::consts::PI / 2.0, std::f32::consts::PI];
let result = neural_ode.forward(&y0, &t_span).unwrap();
// At t = π/2, sin(t) = 1, so dynamics should show this influence
let result_slice = result.to_vec::<f32>().unwrap();
// Basic checks - should have captured time dependence
assert_ne!(result_slice[0], result_slice[2]); // Different values at different times
assert!(result_slice.iter().all(|&x| x.is_finite()));
}
#[test]
fn test_batch_time_series_processing() {
let device = Device::cpu();
let dynamics = TimeSeriesDynamics::new(3, 8, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
// Batch of initial conditions
let y0_batch = Tensor::randn([5, 3], &device).unwrap(); // [batch_size, state_dim]
let t_span = vec![0.0, 0.5, 1.0, 1.5, 2.0];
let result = neural_ode.forward_batch(&y0_batch, &t_span).unwrap();
// Check batch processing output shape
assert_eq!(result.shape(), &[5, 5, 3]); // [batch_size, time_steps, state_dim]
let result_slice = result.to_vec::<f32>().unwrap();
assert!(result_slice.iter().all(|&x| x.is_finite()));
}
}
#[cfg(test)]
mod augmented_neural_odes {
use super::*;
#[test]
fn test_augmented_ode_initialization() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(4, 8, device.clone()).unwrap(); // Original dim = 2, augmented to 4
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let augmented_ode = AugmentedNeuralODE::new(
Box::new(dynamics),
config,
2, // original_dim
2, // augmented_dim
&device,
)
.unwrap();
let y0 = Tensor::randn([2], &device).unwrap(); // Original state
let t_span = vec![0.0, 1.0];
let result = augmented_ode.forward(&y0, &t_span).unwrap();
// Should get back only the original dimensions
assert_eq!(result.shape(), &[2, 2]); // [time_steps, original_dim]
}
#[test]
fn test_augmented_ode_expressivity() {
let device = Device::cpu();
// Create two identical dynamics - one regular, one augmented
let dynamics1 = NeuralNetworkDynamics::new(2, 4, device.clone()).unwrap();
let dynamics2 = NeuralNetworkDynamics::new(4, 8, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let regular_ode = NeuralODE::new(Box::new(dynamics1), config.clone(), &device).unwrap();
let augmented_ode = AugmentedNeuralODE::new(
Box::new(dynamics2),
config,
2, // original_dim
2, // augmented_dim
&device,
)
.unwrap();
let y0 = Tensor::randn([2], &device).unwrap();
let t_span = vec![0.0, 1.0];
let result_regular = regular_ode.forward(&y0, &t_span).unwrap();
let result_augmented = augmented_ode.forward(&y0, &t_span).unwrap();
// Results should have same shape but potentially different values
assert_eq!(result_regular.shape(), result_augmented.shape());
// Both should be finite and reasonable
let reg_slice = result_regular.to_vec::<f32>().unwrap();
let aug_slice = result_augmented.to_vec::<f32>().unwrap();
assert!(reg_slice.iter().all(|&x| x.is_finite()));
assert!(aug_slice.iter().all(|&x| x.is_finite()));
}
}
#[cfg(test)]
mod performance_benchmarks {
use super::*;
use std::time::Instant;
#[test]
fn test_solver_accuracy_vs_speed_tradeoff() {
let device = Device::cpu();
let dynamics = SimpleLinearDynamics::new(device.clone());
let solvers = vec![
(
"Euler",
SolverType::Euler,
SolverConfig::Euler { step_size: 0.1 },
),
(
"RK4",
SolverType::RungeKutta4,
SolverConfig::RungeKutta4 { step_size: 0.1 },
),
(
"Dopri5",
SolverType::Dopri5,
SolverConfig::Dopri5 {
initial_step_size: 0.01,
max_step_size: 0.1,
safety_factor: 0.9,
},
),
];
let y0 = Tensor::from_slice(&[1.0], [1], &device).unwrap();
let t_span = vec![0.0, 1.0];
let expected = (-1.0_f32).exp(); // Analytical solution
let mut results = Vec::new();
for (name, solver_type, solver_config) in solvers {
let config = NeuralODEConfig {
solver: solver_type,
solver_config,
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let start_time = Instant::now();
let result = neural_ode.forward(&y0, &t_span).unwrap();
let elapsed = start_time.elapsed();
let computed_val = result.to_scalar::<f32>().unwrap();
let error = (computed_val - expected).abs();
results.push((name, error, elapsed));
println!("{}: Error = {:.2e}, Time = {:?}", name, error, elapsed);
}
// Verify that we get reasonable accuracy-speed tradeoffs
// Euler should be fastest but least accurate
// Dopri5 should be most accurate but potentially slower
assert!(results.len() == 3);
assert!(results.iter().all(|(_, error, _)| *error < 0.1)); // All should be reasonable
}
#[test]
fn test_memory_scaling_with_time_steps() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(4, 8, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: Some(AdjointConfig {
solver_type: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
rtol: 1e-4,
atol: 1e-7,
}),
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let y0 = tensor_with_grad(Tensor::randn([4], &device).unwrap());
// Test different numbers of time steps
let time_step_counts = vec![5, 10, 20, 50];
for num_steps in time_step_counts {
let t_span: Vec<f32> = (0..num_steps)
.map(|i| i as f32 * 2.0 / (num_steps - 1) as f32)
.collect();
let start_time = Instant::now();
let result = neural_ode.forward(&y0, &t_span).unwrap();
let loss = result.pow_scalar(2.0).unwrap().sum().unwrap();
let _gradients = backward(loss.node_id(), None).unwrap();
let elapsed = start_time.elapsed();
println!("Steps: {}, Time: {:?}", num_steps, elapsed);
// Should complete without memory issues
assert!(elapsed.as_secs() < 30); // Reasonable time limit
}
}
#[test]
fn test_batch_processing_efficiency() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(3, 6, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
let t_span = vec![0.0, 0.5, 1.0];
// Time individual processing
let single_y0 = Tensor::randn([3], &device).unwrap();
let start_single = Instant::now();
for _ in 0..10 {
let _result = neural_ode.forward(&single_y0, &t_span).unwrap();
}
let elapsed_single = start_single.elapsed();
// Time batch processing
let batch_y0 = Tensor::randn([10, 3], &device).unwrap();
let start_batch = Instant::now();
let _result_batch = neural_ode.forward_batch(&batch_y0, &t_span).unwrap();
let elapsed_batch = start_batch.elapsed();
println!(
"Individual: {:?}, Batch: {:?}",
elapsed_single, elapsed_batch
);
// Batch processing should be more efficient (though this is hardware dependent)
// At minimum, both should complete successfully
assert!(elapsed_single.as_millis() > 0);
assert!(elapsed_batch.as_millis() > 0);
}
}
#[cfg(test)]
mod integration_tests {
use super::*;
#[test]
fn test_neural_ode_in_larger_network() {
let device = Device::cpu();
let dynamics = NeuralNetworkDynamics::new(4, 8, device.clone()).unwrap();
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: Some(AdjointConfig {
solver_type: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
rtol: 1e-4,
atol: 1e-7,
}),
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
// Simulate a larger network: input -> linear -> neural_ode -> linear -> output
let input_linear = Tensor::randn([4, 10], &device).unwrap();
let output_linear = Tensor::randn([2, 4], &device).unwrap();
let batch_input = tensor_with_grad(Tensor::randn([5, 10], &device).unwrap());
// Forward pass through larger network
let hidden = batch_input.matmul(&input_linear.t().unwrap()).unwrap();
// Process each sample through Neural ODE
let t_span = vec![0.0, 1.0];
let mut ode_outputs = Vec::new();
for i in 0..5 {
let sample = hidden.slice(0, i..i + 1).unwrap().squeeze(0).unwrap();
let ode_out = neural_ode.forward(&sample, &t_span).unwrap();
let final_state = ode_out.slice(0, -1..-1).unwrap().squeeze(0).unwrap();
ode_outputs.push(final_state);
}
// Stack outputs and continue network
let ode_batch = Tensor::stack(&ode_outputs, 0).unwrap();
let final_output = ode_batch.matmul(&output_linear.t().unwrap()).unwrap();
let loss = final_output.pow_scalar(2.0).unwrap().sum().unwrap();
let gradients = backward(loss.node_id(), None).unwrap();
// Should have gradients for all components
assert!(gradients.contains_key(&batch_input.node_id()));
assert!(!gradients.is_empty());
// Final output should have correct shape
assert_eq!(final_output.shape(), &[5, 2]);
}
#[test]
fn test_neural_ode_reproducibility() {
let device = Device::cpu();
// Create identical dynamics with same random seed
let mut rng = rand::thread_rng();
let seed = rng.r#gen::<u64>();
let dynamics1 = {
use rand::{Rng, SeedableRng};
let mut seeded_rng = rand::rngs::StdRng::seed_from_u64(seed);
NeuralNetworkDynamics::new(2, 4, device.clone()).unwrap()
};
let dynamics2 = {
use rand::{Rng, SeedableRng};
let mut seeded_rng = rand::rngs::StdRng::seed_from_u64(seed);
NeuralNetworkDynamics::new(2, 4, device.clone()).unwrap()
};
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.05 },
adjoint_config: None,
rtol: 1e-6,
atol: 1e-9,
};
let neural_ode1 = NeuralODE::new(Box::new(dynamics1), config.clone(), &device).unwrap();
let neural_ode2 = NeuralODE::new(Box::new(dynamics2), config, &device).unwrap();
let y0 = Tensor::from_slice(&[1.0, -0.5], [2], &device).unwrap();
let t_span = vec![0.0, 0.5, 1.0];
let result1 = neural_ode1.forward(&y0, &t_span).unwrap();
let result2 = neural_ode2.forward(&y0, &t_span).unwrap();
let slice1 = result1.to_vec::<f32>().unwrap();
let slice2 = result2.to_vec::<f32>().unwrap();
// Results should be identical (deterministic computation)
for (v1, v2) in slice1.iter().zip(slice2.iter()) {
assert_relative_eq!(*v1, *v2, epsilon = 1e-8);
}
}
#[test]
fn test_neural_ode_edge_cases() {
let device = Device::cpu();
let dynamics = SimpleLinearDynamics::new(device.clone());
let config = NeuralODEConfig {
solver: SolverType::RungeKutta4,
solver_config: SolverConfig::RungeKutta4 { step_size: 0.1 },
adjoint_config: None,
rtol: 1e-4,
atol: 1e-7,
};
let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap();
// Test edge case: single time point
let y0 = Tensor::from_slice(&[1.0], [1], &device).unwrap();
let t_span_single = vec![0.0];
let result_single = neural_ode.forward(&y0, &t_span_single).unwrap();
assert_eq!(result_single.shape(), &[1, 1]);
// Test edge case: very small time span
let t_span_tiny = vec![0.0, 1e-6];
let result_tiny = neural_ode.forward(&y0, &t_span_tiny).unwrap();
assert_eq!(result_tiny.shape(), &[2, 1]);
let tiny_slice = result_tiny.to_vec::<f32>().unwrap();
assert_relative_eq!(tiny_slice[0], tiny_slice[1], epsilon = 1e-4); // Should be nearly identical
// Test edge case: reverse time integration
let t_span_reverse = vec![1.0, 0.0];
let y1 = Tensor::from_slice(&[(-1.0_f32).exp()], [1], &device).unwrap();
let result_reverse = neural_ode.forward(&y1, &t_span_reverse).unwrap();
let reverse_slice = result_reverse.to_vec::<f32>().unwrap();
assert_relative_eq!(reverse_slice[1], 1.0, epsilon = 0.1); // Should recover initial condition
}
}