use crate::{
optimizer::BaseOptimizer, Optimizer, OptimizerError, OptimizerResult, OptimizerState,
ParamGroup,
};
use parking_lot::RwLock;
use std::collections::HashMap;
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 = 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 = 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_correction2 = 1.0 - self.betas.1.powi(step);
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 = 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 parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
self.base.parameters()
}
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 = 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 = 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 = 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 = 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 parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
self.base.parameters()
}
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()
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
fn const_param(value: f32, shape: &[usize]) -> OptimizerResult<Arc<RwLock<Tensor>>> {
let numel: usize = shape.iter().product();
let tensor = Tensor::from_data(vec![value; numel], shape.to_vec(), DeviceType::Cpu)
.map_err(OptimizerError::TensorError)?;
Ok(Arc::new(RwLock::new(tensor)))
}
fn const_grad(value: f32, shape: &[usize]) -> OptimizerResult<Tensor> {
let numel: usize = shape.iter().product();
Tensor::from_data(vec![value; numel], shape.to_vec(), DeviceType::Cpu)
.map_err(OptimizerError::TensorError)
}
#[test]
fn test_adam_actually_updates_parameters() -> OptimizerResult<()> {
let shape = [4, 4];
let param = const_param(1.0, &shape)?;
let initial = param.read().to_vec().map_err(OptimizerError::TensorError)?;
let mut optimizer = Adam::new(vec![param.clone()], Some(1e-2), None, None, None, false);
for _ in 0..10 {
param.read().set_grad(Some(const_grad(0.5, &shape)?));
optimizer.step()?;
}
let updated = param.read().to_vec().map_err(OptimizerError::TensorError)?;
let total_change: f32 = initial
.iter()
.zip(updated.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
total_change > 1e-4,
"Adam must update parameters: initial={initial:?}, updated={updated:?}, total_change={total_change}"
);
for (a, b) in initial.iter().zip(updated.iter()) {
assert!(
b < a,
"positive gradient must decrease the parameter: {a} -> {b}"
);
}
Ok(())
}
#[test]
fn test_adam_moments_accumulate() -> OptimizerResult<()> {
let shape = [3];
let param = const_param(2.0, &shape)?;
let mut optimizer = Adam::new(vec![param.clone()], Some(1e-3), None, None, None, false);
param.read().set_grad(Some(const_grad(1.0, &shape)?));
optimizer.step()?;
let param_id = format!("{:p}", param.as_ref());
let state = optimizer
.base
.state
.get(¶m_id)
.expect("optimizer state must exist after a step");
let exp_avg = state
.get("exp_avg")
.expect("exp_avg state must exist")
.to_vec()
.map_err(OptimizerError::TensorError)?;
let exp_avg_sq = state
.get("exp_avg_sq")
.expect("exp_avg_sq state must exist")
.to_vec()
.map_err(OptimizerError::TensorError)?;
assert!(
exp_avg.iter().all(|&m| (m - 0.1).abs() < 1e-6),
"first moment must accumulate, got {exp_avg:?}"
);
assert!(
exp_avg_sq.iter().all(|&v| (v - 0.001).abs() < 1e-6),
"second moment must accumulate, got {exp_avg_sq:?}"
);
Ok(())
}
#[test]
fn test_adamw_actually_updates_parameters() -> OptimizerResult<()> {
let shape = [4, 4];
let param = const_param(1.0, &shape)?;
let initial = param.read().to_vec().map_err(OptimizerError::TensorError)?;
let mut optimizer = AdamW::new(
vec![param.clone()],
Some(1e-2),
None,
None,
Some(0.01),
false,
);
for _ in 0..10 {
param.read().set_grad(Some(const_grad(0.5, &shape)?));
optimizer.step()?;
}
let updated = param.read().to_vec().map_err(OptimizerError::TensorError)?;
let total_change: f32 = initial
.iter()
.zip(updated.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
total_change > 1e-4,
"AdamW must update parameters: total_change={total_change}"
);
Ok(())
}
#[test]
fn test_adam_no_grad_no_change() -> OptimizerResult<()> {
let shape = [4];
let param = const_param(1.0, &shape)?;
param.write().set_grad(None);
let before = param.read().to_vec().map_err(OptimizerError::TensorError)?;
let mut optimizer = Adam::new(vec![param.clone()], Some(1e-2), None, None, None, false);
optimizer.step()?;
let after = param.read().to_vec().map_err(OptimizerError::TensorError)?;
assert_eq!(before, after, "Adam must not move params with no gradient");
Ok(())
}
}