980 lines
33 KiB
Rust
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(¶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::<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
|
|
}
|
|
}
|