//! Performance benchmarks for Neural ODE implementations use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; use rtx_transformers::neural_ode::*; use rtx_transformers::prelude::*; fn benchmark_solvers(c: &mut Criterion) { let device = Device::cpu(); let dynamics = ode_func::LinearDynamics::decay(1.0, 4, device.clone()).unwrap(); let solvers = vec![ ( "Euler_0.1", SolverType::Euler, SolverConfig::Euler { step_size: 0.1 }, ), ( "RK4_0.1", 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::ones([4], &device).unwrap(); let t_span = vec![0.0, 1.0]; let mut group = c.benchmark_group("neural_ode_solvers"); for (name, solver_type, solver_config) in solvers { let config = NeuralODEConfig { solver: solver_type, solver_config, adjoint_config: None, rtol: 1e-3, atol: 1e-6, }; let neural_ode = NeuralODE::new( Box::new(ode_func::LinearDynamics::decay(1.0, 4, device.clone()).unwrap()), config, &device, ) .unwrap(); group.bench_with_input(BenchmarkId::new("solver", name), &neural_ode, |b, ode| { b.iter(|| { let _ = ode.forward(&y0, &t_span).unwrap(); }); }); } group.finish(); } fn benchmark_batch_sizes(c: &mut Criterion) { let device = Device::cpu(); let dynamics = ode_func::LinearDynamics::decay(0.5, 2, device.clone()).unwrap(); let config = NeuralODEConfig::default(); let neural_ode = NeuralODE::new(Box::new(dynamics), config, &device).unwrap(); let batch_sizes = vec![1, 10, 50, 100]; let t_span = vec![0.0, 0.5, 1.0]; let mut group = c.benchmark_group("neural_ode_batch_sizes"); for batch_size in batch_sizes { let y0_batch = Tensor::ones([batch_size, 2], &device).unwrap(); group.bench_with_input( BenchmarkId::new("batch_size", batch_size), &y0_batch, |b, y0| { b.iter(|| { let _ = neural_ode.forward_batch(y0, &t_span).unwrap(); }); }, ); } group.finish(); } fn benchmark_cnf(c: &mut Criterion) { let device = Device::cpu(); let dynamics = ode_func::LinearDynamics::decay(0.1, 2, device.clone()).unwrap(); let config = CNFConfig::default(); let cnf = ContinuousNormalizingFlow::new(Box::new(dynamics), config, &device).unwrap(); let z0 = Tensor::randn([50, 2], &device).unwrap(); let t_span = vec![0.0, 1.0]; c.bench_function("cnf_forward", |b| { b.iter(|| { let _ = cnf.forward(&z0, &t_span).unwrap(); }); }); } criterion_group!( benches, benchmark_solvers, benchmark_batch_sizes, benchmark_cnf ); criterion_main!(benches);