use scirs2_core::ndarray::{s, Array1, Array2};
use scirs2_core::random::distributions::Uniform;
use scirs2_core::random::RandomExt;
use optirs_core::optimizers::{Adam, Lookahead, SGD};
use optirs_core::Optimizer;
use std::error::Error;
use std::time::Instant;
#[allow(dead_code)]
fn compute_loss_and_gradient(
weights: &Array1<f64>,
bias: &f64,
x: &Array2<f64>,
y: &Array1<f64>,
) -> (f64, Array1<f64>, f64) {
let preds = x.dot(weights) + *bias;
let diff = &preds - y;
let loss = diff.mapv(|v| v * v).mean().expect("unwrap failed");
let n = x.nrows() as f64;
let grad_w = x.t().dot(&diff) * (2.0 / n);
let grad_b = diff.sum() * (2.0 / n);
(loss, grad_w, grad_b)
}
#[allow(dead_code)]
fn train_model(
optimizer: &mut dyn Optimizer<f64, scirs2_core::ndarray::Ix1>,
x: &Array2<f64>,
y: &Array1<f64>,
train_batches: &Vec<(Array2<f64>, Array1<f64>)>,
lr: f64,
epochs: usize,
) -> Result<(Array1<f64>, f64, f64), Box<dyn Error>> {
let n_features = x.ncols();
let mut weights = Array1::zeros(n_features);
let mut bias = 0.0;
let start_time = Instant::now();
let n_batches = train_batches.len();
for epoch in 0..epochs {
let mut epoch_loss = 0.0;
for (batch_x, batch_y) in train_batches {
let (loss, grad_w, grad_b) =
compute_loss_and_gradient(&weights, &bias, batch_x, batch_y);
epoch_loss += loss;
weights = optimizer.step(&weights, &grad_w)?;
bias -= lr * grad_b;
}
epoch_loss /= n_batches as f64;
if epoch % 5 == 0 || epoch == epochs - 1 {
println!(" Epoch {}: Loss = {:.6}", epoch, epoch_loss);
}
}
let _duration = start_time.elapsed();
let (final_loss, _, _) = compute_loss_and_gradient(&weights, &bias, x, y);
Ok((weights, bias, final_loss))
}
#[allow(dead_code)]
fn main() -> Result<(), Box<dyn Error>> {
let n_samples = 1000;
let n_features = 20;
println!("Comparing optimizers on noisy linear regression task");
println!("====================================================");
println!("Dataset: {} samples, {} features", n_samples, n_features);
println!();
let x = Array2::random((n_samples, n_features), Uniform::new(-1.0, 1.0));
let true_weights = Array1::random(n_features, Uniform::new(-1.0, 1.0));
let true_bias = 0.5;
let noise = Array1::random(n_samples, Uniform::new(-0.3, 0.3)); let y = x.dot(&true_weights) + true_bias + noise;
let batch_size = 50;
let n_batches = n_samples / batch_size;
let train_batches: Vec<_> = (0..n_batches)
.map(|i| {
let start = i * batch_size;
let end = start + batch_size;
let x_batch = x.slice(s![start..end, ..]).to_owned();
let y_batch = y.slice(s![start..end]).to_owned();
(x_batch, y_batch)
})
.collect();
let lr = 0.01;
let epochs = 30;
println!("Training with SGD:");
let mut sgd = SGD::new(lr).with_momentum(0.9).with_weight_decay(0.0001);
let (sgd_weights, sgd_bias, sgd_loss) =
train_model(&mut sgd, &x, &y, &train_batches, lr, epochs)?;
println!("\nTraining with Adam:");
let mut adam = Adam::new(lr).with_weight_decay(0.0001);
let (adam_weights, adam_bias, adam_loss) =
train_model(&mut adam, &x, &y, &train_batches, lr, epochs)?;
println!("\nTraining with Lookahead(SGD):");
let sgd = SGD::new(lr).with_momentum(0.9).with_weight_decay(0.0001);
let mut lookahead_sgd = Lookahead::with_config(sgd, 0.5, 5);
let (lookahead_sgd_weights, lookahead_sgd_bias, lookahead_sgd_loss) =
train_model(&mut lookahead_sgd, &x, &y, &train_batches, lr, epochs)?;
println!("\nTraining with Lookahead(Adam):");
let adam = Adam::new(lr).with_weight_decay(0.0001);
let mut lookahead_adam = Lookahead::with_config(adam, 0.5, 5);
let (lookahead_adam_weights, lookahead_adam_bias, lookahead_adam_loss) =
train_model(&mut lookahead_adam, &x, &y, &train_batches, lr, epochs)?;
println!("\nResults Comparison:");
println!("===================");
let sgd_weight_error = (&sgd_weights - &true_weights)
.mapv(|v| v.abs())
.mean()
.expect("unwrap failed");
let sgd_bias_error = (sgd_bias - true_bias).abs();
let adam_weight_error = (&adam_weights - &true_weights)
.mapv(|v| v.abs())
.mean()
.expect("unwrap failed");
let adam_bias_error = (adam_bias - true_bias).abs();
let lookahead_sgd_weight_error = (&lookahead_sgd_weights - &true_weights)
.mapv(|v| v.abs())
.mean()
.expect("unwrap failed");
let lookahead_sgd_bias_error = (lookahead_sgd_bias - true_bias).abs();
let lookahead_adam_weight_error = (&lookahead_adam_weights - &true_weights)
.mapv(|v| v.abs())
.mean()
.expect("unwrap failed");
let lookahead_adam_bias_error = (lookahead_adam_bias - true_bias).abs();
println!(
"{:<17} {:<10} {:<12} {:<10}",
"Optimizer", "Loss", "Weight Error", "Bias Error"
);
println!("{:-<17} {:-<10} {:-<12} {:-<10}", "", "", "", "");
println!(
"{:<17} {:<10.6} {:<12.6} {:<10.6}",
"SGD", sgd_loss, sgd_weight_error, sgd_bias_error
);
println!(
"{:<17} {:<10.6} {:<12.6} {:<10.6}",
"Adam", adam_loss, adam_weight_error, adam_bias_error
);
println!(
"{:<17} {:<10.6} {:<12.6} {:<10.6}",
"Lookahead(SGD)", lookahead_sgd_loss, lookahead_sgd_weight_error, lookahead_sgd_bias_error
);
println!(
"{:<17} {:<10.6} {:<12.6} {:<10.6}",
"Lookahead(Adam)",
lookahead_adam_loss,
lookahead_adam_weight_error,
lookahead_adam_bias_error
);
println!("\nConclusions:");
println!("============");
println!("Lookahead typically provides more stable optimization by taking");
println!("\"k steps forward, 1 step back,\" which can lead to better generalization");
println!("performance. It is particularly effective when used with SGD, often");
println!("matching or exceeding the performance of Adam while maintaining SGD's");
println!("generalization benefits. This comes at minimal computational overhead.");
Ok(())
}