use optirs_core::{
optimizers::{Optimizer, SGD},
schedulers::{CombinedScheduler, CustomScheduler, LearningRateScheduler, SchedulerBuilder},
};
use scirs2_core::ndarray::Array1;
use std::error::Error;
#[allow(dead_code)]
fn quadratic_loss(x: &Array1<f64>) -> f64 {
x.iter().map(|&xi| xi * xi).sum()
}
#[allow(dead_code)]
fn quadratic_gradient(x: &Array1<f64>) -> Array1<f64> {
x * 2.0
}
#[allow(dead_code)]
fn test_scheduler<S: LearningRateScheduler<f64>>(
name: &str,
mut scheduler: S,
steps: usize,
initial_params: &Array1<f64>,
) -> Result<Array1<f64>, Box<dyn Error>> {
println!("\n{}", "=".repeat(50));
println!("Testing scheduler: {}", name);
println!("{}", "=".repeat(50));
println!("Learning rate schedule:");
println!("Step | Learning Rate");
println!("-----|-------------");
let mut params = initial_params.clone();
let mut optimizer = SGD::<f64>::new(scheduler.get_learning_rate());
for step in 0..steps {
if step % (steps / 10).max(1) == 0 || step == steps - 1 {
println!("{:4} | {:13.8}", step, scheduler.get_learning_rate());
}
<SGD<f64> as Optimizer<f64, scirs2_core::ndarray::Ix1>>::set_learning_rate(
&mut optimizer,
scheduler.get_learning_rate(),
);
let gradient = quadratic_gradient(¶ms);
params = optimizer.step(¶ms, &gradient)?;
if step < steps - 1 {
scheduler.step();
}
}
println!("\nFinal loss: {:.8}", quadratic_loss(¶ms));
println!("Final parameters: {:?}", params);
Ok(params)
}
#[allow(dead_code)]
fn main() -> Result<(), Box<dyn Error>> {
let initial_params = Array1::from_vec(vec![5.0, -3.0, 2.0, -4.0]);
let initial_loss = quadratic_loss(&initial_params);
println!("Initial loss: {:.8}", initial_loss);
println!("Initial parameters: {:?}", initial_params);
let steps = 100;
let custom_func_scheduler = CustomScheduler::new(0.1, |step| {
let step = step as f64;
let warmup_steps = 10.0;
let decay_rate: f64 = 0.9;
if step < warmup_steps {
0.01 + (0.1 - 0.01) * (step / warmup_steps)
} else {
0.1 * decay_rate.powf((step - warmup_steps) / 10.0)
}
});
let params1 = test_scheduler(
"Custom Function Scheduler",
custom_func_scheduler,
steps,
&initial_params,
)?;
let step_scheduler = SchedulerBuilder::new(0.1).step_decay(20, 0.5);
let params2 = test_scheduler(
"Step Decay Scheduler (from builder)",
step_scheduler,
steps,
&initial_params,
)?;
let exp_scheduler = CustomScheduler::new(0.1, |step| 0.1 * 0.95f64.powi((step / 5) as i32));
let cosine_scheduler = SchedulerBuilder::new(0.1).cosine_annealing(steps, 0.001);
let combined_scheduler = CombinedScheduler::new(
exp_scheduler,
cosine_scheduler,
|lr1, lr2| lr1.min(lr2),
);
let params3 = test_scheduler(
"Combined Scheduler (Exp + Cosine)",
combined_scheduler,
steps,
&initial_params,
)?;
let staircase_scheduler = CustomScheduler::new(0.1, move |step| {
let stair_size = 20;
let num_stairs = (step / stair_size) as i32;
0.1 * 0.5f64.powi(num_stairs)
});
let params4 = test_scheduler(
"Staircase Scheduler",
staircase_scheduler,
steps,
&initial_params,
)?;
let mut rng = scirs2_core::random::thread_rng();
let noisy_scheduler = CustomScheduler::new(0.1, move |step| {
let base_lr = 0.1 * 0.95f64.powi((step / 10) as i32);
let noise = rng.gen_range(-0.01..0.01); (base_lr + noise).max(0.001) });
let params5 = test_scheduler("Noisy Scheduler", noisy_scheduler, steps, &initial_params)?;
println!("\n{}", "=".repeat(50));
println!("Final Loss Comparison");
println!("{}", "=".repeat(50));
println!("Initial loss: {:.8}", initial_loss);
println!("Custom Function Scheduler: {:.8}", quadratic_loss(¶ms1));
println!("Step Decay Scheduler: {:.8}", quadratic_loss(¶ms2));
println!("Combined Scheduler: {:.8}", quadratic_loss(¶ms3));
println!("Staircase Scheduler: {:.8}", quadratic_loss(¶ms4));
println!("Noisy Scheduler: {:.8}", quadratic_loss(¶ms5));
Ok(())
}