//! 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 { // dx/dt = -x (simple exponential decay) y.neg() } fn parameters(&self) -> Vec { 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 { 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 { // 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 { 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::().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::().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::().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::().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(¶m.node_id())); let grad = &gradients[¶m.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::() .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 { 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 { let h = y.matmul(&self.linear.t()?)?.add(&self.bias)?; h.tanh() // Bounded dynamics for stability } fn parameters(&self) -> Vec { 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::().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::().unwrap(); let x0_recon_slice = x0_reconstructed.to_vec::().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::().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, device: Device, } impl TimeSeriesDynamics { fn new(input_dim: usize, hidden_dim: usize, device: Device) -> Result { 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 { // 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 { 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::().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::().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::().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::().unwrap(); let aug_slice = result_augmented.to_vec::().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::().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 = (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::(); 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::().unwrap(); let slice2 = result2.to_vec::().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::().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::().unwrap(); assert_relative_eq!(reverse_slice[1], 1.0, epsilon = 0.1); // Should recover initial condition } }