use parking_lot::RwLock;
use std::sync::Arc;
use std::time::Instant;
use torsh_core::error::Result;
use torsh_optim::prelude::*;
use torsh_tensor::Tensor;
#[derive(Debug, Clone)]
struct BenchmarkResult {
optimizer_name: String,
final_loss: f32,
iterations: usize,
total_time_ms: f64,
time_per_step_us: f64,
converged: bool,
}
impl BenchmarkResult {
fn print(&self) {
println!(
" {:20} | Loss: {:8.6} | Steps: {:4} | Time: {:6.2}ms | Per-step: {:6.2}ยตs | {}",
self.optimizer_name,
self.final_loss,
self.iterations,
self.total_time_ms,
self.time_per_step_us,
if self.converged { "โ" } else { "โ" }
);
}
}
fn benchmark_optimizer<O: Optimizer>(
name: &str,
mut optimizer: O,
param: Arc<RwLock<Tensor>>,
max_iterations: usize,
convergence_threshold: f32,
) -> Result<BenchmarkResult> {
let start = Instant::now();
let mut iterations = 0;
let mut final_loss = f32::INFINITY;
let mut converged = false;
for i in 0..max_iterations {
let grad = { param.read().mul_scalar(2.0)? };
param.write().set_grad(Some(grad));
optimizer.step().map_err(|e| {
torsh_core::error::TorshError::RuntimeError(format!("Step failed: {:?}", e))
})?;
let loss = {
let val = param.read().to_vec()?[0];
val * val
};
final_loss = loss;
iterations = i + 1;
optimizer.zero_grad();
if loss < convergence_threshold {
converged = true;
break;
}
}
let elapsed = start.elapsed();
let total_time_ms = elapsed.as_secs_f64() * 1000.0;
let time_per_step_us = (total_time_ms * 1000.0) / iterations as f64;
Ok(BenchmarkResult {
optimizer_name: name.to_string(),
final_loss,
iterations,
total_time_ms,
time_per_step_us,
converged,
})
}
fn main() -> Result<()> {
println!("โก Optimizer Performance Benchmark\n");
println!("{}", "=".repeat(90));
println!("\n๐ Problem: Minimize f(x) = xยฒ starting from x=10.0");
println!("๐ฏ Target: Loss < 1e-6");
println!("โฑ๏ธ Max iterations: 1000\n");
let max_iterations = 1000;
let convergence_threshold = 1e-6;
let initial_value = 10.0;
let mut results = Vec::new();
println!("\n{}", "=".repeat(90));
println!("๐ Category 1: First-Order Optimizers");
println!("{}", "=".repeat(90));
println!(
"\n {:20} | {:^14} | {:^10} | {:^12} | {:^14} | Status",
"Optimizer", "Final Loss", "Iterations", "Total Time", "Time/Step"
);
println!(" {}", "-".repeat(88));
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = SGD::new(vec![param.clone()], 0.1, Some(0.9), None, None, false);
let result = benchmark_optimizer(
"SGD",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = Adam::new(
vec![param.clone()],
Some(0.1),
Some((0.9, 0.999)),
None,
None,
false,
);
let result = benchmark_optimizer(
"Adam",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = AdaGrad::new(
vec![param.clone()],
Some(1.0), Some(0.0), Some(0.0), None, Some(1e-10), );
let result = benchmark_optimizer(
"AdaGrad",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = RMSprop::new(
vec![param.clone()],
Some(0.01), Some(0.99), Some(1e-8), Some(0.0), Some(0.0), false, );
let result = benchmark_optimizer(
"RMSprop",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
println!("\n\n{}", "=".repeat(90));
println!("๐ Category 2: Advanced Adaptive Methods");
println!("{}", "=".repeat(90));
println!(
"\n {:20} | {:^14} | {:^10} | {:^12} | {:^14} | Status",
"Optimizer", "Final Loss", "Iterations", "Total Time", "Time/Step"
);
println!(" {}", "-".repeat(88));
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = RAdam::new(
vec![param.clone()],
Some(0.1),
Some(0.9),
Some(0.999),
Some(1e-8),
None,
);
let result = benchmark_optimizer(
"RAdam",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = AdaBelief::new(vec![param.clone()], 0.1); let result = benchmark_optimizer(
"AdaBelief",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = NAdam::new(vec![param.clone()], 0.1); let result = benchmark_optimizer(
"NAdam",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
println!("\n\n{}", "=".repeat(90));
println!("๐ Category 3: Modern Optimizers (2023-2024)");
println!("{}", "=".repeat(90));
println!(
"\n {:20} | {:^14} | {:^10} | {:^12} | {:^14} | Status",
"Optimizer", "Final Loss", "Iterations", "Total Time", "Time/Step"
);
println!(" {}", "-".repeat(88));
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = Lion::builder()
.params(vec![param.clone()])
.lr(0.01)
.beta1(0.9)
.beta2(0.99)
.build();
let result = benchmark_optimizer(
"Lion",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let optimizer = Prodigy::builder()
.params(vec![param.clone()])
.lr(1.0)
.beta1(0.9)
.beta2(0.999)
.build();
let result = benchmark_optimizer(
"Prodigy",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
println!("\n\n{}", "=".repeat(90));
println!("๐ Category 4: Meta-Optimizers");
println!("{}", "=".repeat(90));
println!(
"\n {:20} | {:^14} | {:^10} | {:^12} | {:^14} | Status",
"Optimizer", "Final Loss", "Iterations", "Total Time", "Time/Step"
);
println!(" {}", "-".repeat(88));
{
let param = Arc::new(RwLock::new(Tensor::scalar(initial_value)?));
let base = Adam::new(
vec![param.clone()],
Some(0.1),
Some((0.9, 0.999)),
None,
None,
false,
);
let optimizer = Lookahead::new(base, 0.5, 5); let result = benchmark_optimizer(
"Lookahead(Adam)",
optimizer,
param,
max_iterations,
convergence_threshold,
)?;
result.print();
results.push(result);
}
println!("\n\n{}", "=".repeat(90));
println!("๐ Performance Analysis");
println!("{}", "=".repeat(90));
let fastest = results
.iter()
.filter(|r| r.converged)
.min_by_key(|r| r.iterations);
if let Some(fastest) = fastest {
println!("\n๐ Fastest Convergence: {}", fastest.optimizer_name);
println!(" Converged in {} iterations", fastest.iterations);
println!(" Final loss: {:.6}", fastest.final_loss);
}
let best_loss = results.iter().min_by(|a, b| {
a.final_loss
.partial_cmp(&b.final_loss)
.unwrap_or(std::cmp::Ordering::Equal)
});
if let Some(best) = best_loss {
println!("\n๐ฏ Best Final Loss: {}", best.optimizer_name);
println!(" Final loss: {:.6}", best.final_loss);
println!(" Iterations: {}", best.iterations);
}
let fastest_step = results.iter().min_by(|a, b| {
a.time_per_step_us
.partial_cmp(&b.time_per_step_us)
.unwrap_or(std::cmp::Ordering::Equal)
});
if let Some(fastest) = fastest_step {
println!("\nโก Fastest Per-Step: {}", fastest.optimizer_name);
println!(" Time per step: {:.2}ยตs", fastest.time_per_step_us);
}
let converged_count = results.iter().filter(|r| r.converged).count();
println!("\n๐ Convergence Summary:");
println!(
" {}/{} optimizers converged to threshold",
converged_count,
results.len()
);
println!("\n\n{}", "=".repeat(90));
println!("๐ก Recommendations Based on Benchmarks");
println!("{}", "=".repeat(90));
println!("\n๐ฏ For Fast Convergence:");
println!(" โ Adam, RAdam, or NAdam");
println!(" โ Reliable, fast, and work out-of-the-box");
println!("\n๐พ For Memory Efficiency:");
println!(" โ Lion (only momentum, no second moment)");
println!(" โ SGD with momentum (minimal state)");
println!("\n๐ง For No Hyperparameter Tuning:");
println!(" โ Prodigy (auto-adaptive learning rate)");
println!(" โ Just use lr=1.0 and go!");
println!("\nโก For Speed:");
println!(" โ SGD (simplest, fastest per-step)");
println!(" โ But may need more iterations to converge");
println!("\n๐ For Best Final Performance:");
println!(" โ Often depends on the problem");
println!(" โ Try multiple optimizers with cross-validation");
println!("\n\n{}", "=".repeat(90));
println!("โ
Benchmark completed successfully!");
println!("๐ก Run with --release flag for accurate timing measurements");
println!("{}", "=".repeat(90));
Ok(())
}