use crate::{
optimizer::BaseOptimizer, Optimizer, OptimizerError, OptimizerResult, OptimizerState,
ParamGroup,
};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::ops::Add;
use std::sync::Arc;
use torsh_core::error::Result;
use torsh_tensor::{creation::zeros_like, Tensor};
#[derive(Clone)]
pub struct Adam {
base: BaseOptimizer,
betas: (f32, f32),
eps: f32,
weight_decay: f32,
amsgrad: bool,
}
impl Adam {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: Option<f32>,
betas: Option<(f32, f32)>,
eps: Option<f32>,
weight_decay: Option<f32>,
amsgrad: bool,
) -> Self {
let lr = lr.unwrap_or(1e-3);
let betas = betas.unwrap_or((0.9, 0.999));
let eps = eps.unwrap_or(1e-8);
let weight_decay = weight_decay.unwrap_or(0.0);
let mut defaults = HashMap::new();
defaults.insert("lr".to_string(), lr);
defaults.insert("beta1".to_string(), betas.0);
defaults.insert("beta2".to_string(), betas.1);
defaults.insert("eps".to_string(), eps);
defaults.insert("weight_decay".to_string(), weight_decay);
let param_group = ParamGroup::new(params, lr);
let base = BaseOptimizer {
param_groups: vec![param_group],
state: HashMap::new(),
optimizer_type: "Adam".to_string(),
defaults,
};
Self {
base,
betas,
eps,
weight_decay,
amsgrad,
}
}
}
impl Optimizer for Adam {
fn step(&mut self) -> OptimizerResult<()> {
for group in &mut self.base.param_groups {
for param_arc in &group.params {
let mut param = param_arc.write();
if !param.has_grad() {
continue;
}
let grad = param
.grad()
.expect("gradient should exist after has_grad check");
let param_id = format!("{:p}", param_arc.as_ref());
let needs_init = !self.base.state.contains_key(¶m_id);
let state = self
.base
.state
.entry(param_id.clone())
.or_insert_with(HashMap::new);
if needs_init {
state.insert("step".to_string(), zeros_like(¶m)?);
state.insert("exp_avg".to_string(), zeros_like(¶m)?);
state.insert("exp_avg_sq".to_string(), zeros_like(¶m)?);
if self.amsgrad {
state.insert("max_exp_avg_sq".to_string(), zeros_like(¶m)?);
}
}
let mut step_tensor = state.get("step").expect("step state should exist").clone();
let mut exp_avg = state
.get("exp_avg")
.expect("exp_avg state should exist")
.clone();
let mut exp_avg_sq = state
.get("exp_avg_sq")
.expect("exp_avg_sq state should exist")
.clone();
step_tensor
.add_scalar_(1.0)
.map_err(OptimizerError::TensorError)?;
let step = step_tensor.to_vec().map_err(OptimizerError::TensorError)?[0] as i32;
let mut grad = grad;
if self.weight_decay != 0.0 {
let weight_decay_term = param
.mul_scalar(self.weight_decay)
.map_err(OptimizerError::TensorError)?;
grad = grad
.add(&weight_decay_term)
.map_err(OptimizerError::TensorError)?;
}
exp_avg
.mul_scalar_(self.betas.0)
.map_err(OptimizerError::TensorError)?;
let grad_term = grad
.mul_scalar(1.0 - self.betas.0)
.map_err(OptimizerError::TensorError)?;
exp_avg
.add(&grad_term)
.map_err(OptimizerError::TensorError)?;
exp_avg_sq
.mul_scalar_(self.betas.1)
.map_err(OptimizerError::TensorError)?;
let grad_squared = grad.mul_op(&grad).map_err(OptimizerError::TensorError)?;
let grad_sq_term = grad_squared
.mul_scalar(1.0 - self.betas.1)
.map_err(OptimizerError::TensorError)?;
exp_avg_sq
.add(&grad_sq_term)
.map_err(OptimizerError::TensorError)?;
let denom = if self.amsgrad {
let mut max_exp_avg_sq = state
.get("max_exp_avg_sq")
.expect("max_exp_avg_sq state should exist")
.clone();
max_exp_avg_sq = max_exp_avg_sq
.maximum(&exp_avg_sq)
.map_err(OptimizerError::TensorError)?;
state.insert("max_exp_avg_sq".to_string(), max_exp_avg_sq.clone());
let sqrt_max = max_exp_avg_sq.sqrt().map_err(OptimizerError::TensorError)?;
sqrt_max
.add_scalar(self.eps)
.map_err(OptimizerError::TensorError)?
} else {
let bias_correction1 = 1.0 - self.betas.0.powi(step);
let bias_correction2 = 1.0 - self.betas.1.powi(step);
let _corrected_exp_avg = exp_avg
.div_scalar(bias_correction1)
.map_err(OptimizerError::TensorError)?;
let corrected_exp_avg_sq = exp_avg_sq
.div_scalar(bias_correction2)
.map_err(OptimizerError::TensorError)?;
let sqrt_corrected = corrected_exp_avg_sq
.sqrt()
.map_err(OptimizerError::TensorError)?;
sqrt_corrected
.add_scalar(self.eps)
.map_err(OptimizerError::TensorError)?
};
let step_size = group.lr;
let bias_correction1 = 1.0 - self.betas.0.powi(step);
let corrected_exp_avg = exp_avg
.div_scalar(bias_correction1)
.map_err(OptimizerError::TensorError)?;
let update = corrected_exp_avg
.div(&denom)
.map_err(OptimizerError::TensorError)?
.mul_scalar(step_size)
.map_err(OptimizerError::TensorError)?;
param.sub(&update).map_err(OptimizerError::TensorError)?;
state.insert("step".to_string(), step_tensor);
state.insert("exp_avg".to_string(), exp_avg);
state.insert("exp_avg_sq".to_string(), exp_avg_sq);
}
}
Ok(())
}
fn zero_grad(&mut self) {
self.base.zero_grad();
}
fn get_lr(&self) -> Vec<f32> {
self.base.get_lr()
}
fn set_lr(&mut self, lr: f32) {
self.base.set_lr(lr);
}
fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
self.base.add_param_group(params, options);
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
self.base.state_dict()
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
self.base.load_state_dict(state)
}
}
pub struct AdamW {
base: BaseOptimizer,
betas: (f32, f32),
eps: f32,
weight_decay: f32,
amsgrad: bool,
}
impl AdamW {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: Option<f32>,
betas: Option<(f32, f32)>,
eps: Option<f32>,
weight_decay: Option<f32>,
amsgrad: bool,
) -> Self {
let lr = lr.unwrap_or(1e-3);
let betas = betas.unwrap_or((0.9, 0.999));
let eps = eps.unwrap_or(1e-8);
let weight_decay = weight_decay.unwrap_or(0.01);
let mut defaults = HashMap::new();
defaults.insert("lr".to_string(), lr);
defaults.insert("beta1".to_string(), betas.0);
defaults.insert("beta2".to_string(), betas.1);
defaults.insert("eps".to_string(), eps);
defaults.insert("weight_decay".to_string(), weight_decay);
let param_group = ParamGroup::new(params, lr);
let base = BaseOptimizer {
param_groups: vec![param_group],
state: HashMap::new(),
optimizer_type: "AdamW".to_string(),
defaults,
};
Self {
base,
betas,
eps,
weight_decay,
amsgrad,
}
}
}
impl Optimizer for AdamW {
fn step(&mut self) -> OptimizerResult<()> {
for group in &mut self.base.param_groups {
for param_arc in &group.params {
let mut param = param_arc.write();
if !param.has_grad() {
continue;
}
let grad = param
.grad()
.expect("gradient should exist after has_grad check");
let param_id = format!("{:p}", param_arc.as_ref());
let needs_init = !self.base.state.contains_key(¶m_id);
let state = self
.base
.state
.entry(param_id.clone())
.or_insert_with(HashMap::new);
if needs_init {
state.insert("step".to_string(), zeros_like(¶m)?);
state.insert("exp_avg".to_string(), zeros_like(¶m)?);
state.insert("exp_avg_sq".to_string(), zeros_like(¶m)?);
if self.amsgrad {
state.insert("max_exp_avg_sq".to_string(), zeros_like(¶m)?);
}
}
let mut step_tensor = state.get("step").expect("step state should exist").clone();
let mut exp_avg = state
.get("exp_avg")
.expect("exp_avg state should exist")
.clone();
let mut exp_avg_sq = state
.get("exp_avg_sq")
.expect("exp_avg_sq state should exist")
.clone();
step_tensor
.add_scalar_(1.0)
.map_err(OptimizerError::TensorError)?;
let step = step_tensor.to_vec().map_err(OptimizerError::TensorError)?[0] as i32;
if self.weight_decay != 0.0 {
let weight_decay_update = param
.mul_scalar(group.lr * self.weight_decay)
.map_err(OptimizerError::TensorError)?;
param
.sub(&weight_decay_update)
.map_err(OptimizerError::TensorError)?;
}
exp_avg
.mul_scalar_(self.betas.0)
.map_err(OptimizerError::TensorError)?;
let grad_term = grad
.mul_scalar(1.0 - self.betas.0)
.map_err(OptimizerError::TensorError)?;
exp_avg
.add(&grad_term)
.map_err(OptimizerError::TensorError)?;
exp_avg_sq
.mul_scalar_(self.betas.1)
.map_err(OptimizerError::TensorError)?;
let grad_squared = grad.mul_op(&grad).map_err(OptimizerError::TensorError)?;
let grad_sq_term = grad_squared
.mul_scalar(1.0 - self.betas.1)
.map_err(OptimizerError::TensorError)?;
exp_avg_sq
.add(&grad_sq_term)
.map_err(OptimizerError::TensorError)?;
let bias_correction1 = 1.0 - self.betas.0.powi(step);
let bias_correction2 = 1.0 - self.betas.1.powi(step);
let corrected_exp_avg = exp_avg
.div_scalar(bias_correction1)
.map_err(OptimizerError::TensorError)?;
let corrected_exp_avg_sq = exp_avg_sq
.div_scalar(bias_correction2)
.map_err(OptimizerError::TensorError)?;
let denom = if self.amsgrad {
let mut max_exp_avg_sq = state
.get("max_exp_avg_sq")
.expect("max_exp_avg_sq state should exist")
.clone();
max_exp_avg_sq = max_exp_avg_sq
.maximum(&corrected_exp_avg_sq)
.map_err(OptimizerError::TensorError)?;
state.insert("max_exp_avg_sq".to_string(), max_exp_avg_sq.clone());
let sqrt_max = max_exp_avg_sq.sqrt().map_err(OptimizerError::TensorError)?;
sqrt_max
.add_scalar(self.eps)
.map_err(OptimizerError::TensorError)?
} else {
let sqrt_corrected = corrected_exp_avg_sq
.sqrt()
.map_err(OptimizerError::TensorError)?;
sqrt_corrected
.add_scalar(self.eps)
.map_err(OptimizerError::TensorError)?
};
let step_size = group.lr;
let update = corrected_exp_avg
.div(&denom)
.map_err(OptimizerError::TensorError)?
.mul_scalar(step_size)
.map_err(OptimizerError::TensorError)?;
param.sub(&update).map_err(OptimizerError::TensorError)?;
state.insert("step".to_string(), step_tensor);
state.insert("exp_avg".to_string(), exp_avg);
state.insert("exp_avg_sq".to_string(), exp_avg_sq);
}
}
Ok(())
}
fn zero_grad(&mut self) {
self.base.zero_grad();
}
fn get_lr(&self) -> Vec<f32> {
self.base.get_lr()
}
fn set_lr(&mut self, lr: f32) {
self.base.set_lr(lr);
}
fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
self.base.add_param_group(params, options);
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
self.base.state_dict()
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
self.base.load_state_dict(state)
}
}
pub struct AdamBuilder {
lr: f32,
betas: (f32, f32),
eps: f32,
weight_decay: f32,
amsgrad: bool,
}
impl AdamBuilder {
pub fn new() -> Self {
Self {
lr: 1e-3,
betas: (0.9, 0.999),
eps: 1e-8,
weight_decay: 0.0,
amsgrad: false,
}
}
pub fn lr(mut self, lr: f32) -> Self {
self.lr = lr;
self
}
pub fn betas(mut self, beta1: f32, beta2: f32) -> Self {
self.betas = (beta1, beta2);
self
}
pub fn eps(mut self, eps: f32) -> Self {
self.eps = eps;
self
}
pub fn weight_decay(mut self, weight_decay: f32) -> Self {
self.weight_decay = weight_decay;
self
}
pub fn amsgrad(mut self, amsgrad: bool) -> Self {
self.amsgrad = amsgrad;
self
}
pub fn build(self, params: Vec<Arc<RwLock<Tensor>>>) -> Adam {
Adam::new(
params,
Some(self.lr),
Some(self.betas),
Some(self.eps),
Some(self.weight_decay),
self.amsgrad,
)
}
pub fn build_adamw(self, params: Vec<Arc<RwLock<Tensor>>>) -> AdamW {
AdamW::new(
params,
Some(self.lr),
Some(self.betas),
Some(self.eps),
Some(self.weight_decay),
self.amsgrad,
)
}
}
impl Default for AdamBuilder {
fn default() -> Self {
Self::new()
}
}