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

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);