115 lines
3.2 KiB
Rust
115 lines
3.2 KiB
Rust
//! 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);
|