use candle_core::{backprop::GradStore, Result as CandleResult, Var};
pub use crate::nn::optim::adam::Adam;
pub use crate::nn::optim::rmsprop::RMSprop;
pub use crate::nn::optim::sgd::SGD;
mod adam;
mod rmsprop;
mod sgd;
#[derive(Clone, Copy, Debug, PartialEq, PartialOrd)]
pub enum Decay {
WeightDecay(f64),
DecoupledWeightDecay(f64),
}
#[derive(Copy, Clone, Debug, PartialEq, PartialOrd)]
pub enum Momentum {
Classical(f64),
Nesterov(f64),
}
pub trait Optimizer {
fn step(&mut self, grads: &GradStore, vars: Vec<&mut Var>) -> CandleResult<()>;
fn learning_rate(&self) -> f64;
fn set_learning_rate(&mut self, lr: f64);
}
pub enum OptimizerParams {
SGD {
dampening: Option<f64>,
weight_decay: Option<Decay>,
momentum: Option<Momentum>,
},
AdamW {
beta_1: f64,
beta_2: f64,
eps: f64,
weight_decay: Option<Decay>,
amsgrad: bool,
},
Else,
}
#[derive(Debug, Clone)]
pub enum OptimKind {
SGD,
Adam,
AdamW, RMSprop,
}
pub fn create_optimizer<'a>(
kind: OptimKind,
vars: Vec<&mut Var>,
lr: f64,
params: OptimizerParams,
) -> CandleResult<Box<dyn Optimizer + 'a>> {
match kind {
OptimKind::SGD => {
if let OptimizerParams::SGD {
momentum,
weight_decay,
..
} = params
{
let mut opt = SGD::new(lr)?;
if let Some(Momentum::Nesterov(val)) = momentum {
opt = opt.momentum(val);
}
if let Some(Decay::WeightDecay(val)) = weight_decay {
opt = opt.weight_decay(val);
}
Ok(Box::new(opt))
} else {
Ok(Box::new(SGD::new(lr)?))
}
}
OptimKind::Adam | OptimKind::AdamW => {
if let OptimizerParams::AdamW {
beta_1,
beta_2,
eps,
weight_decay,
..
} = params
{
let mut opt = Adam::new(vars, lr)?.betas(beta_1, beta_2).eps(eps);
if let Some(wd) = weight_decay {
match wd {
Decay::WeightDecay(val) => {
opt = opt.weight_decay(val);
}
Decay::DecoupledWeightDecay(val) => {
opt = opt.decoupled_weight_decay(val);
}
}
}
Ok(Box::new(opt))
} else {
Ok(Box::new(Adam::new(vars, lr)?))
}
}
OptimKind::RMSprop => Ok(Box::new(RMSprop::new(vars, lr)?)),
}
}