use crate::{
Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use torsh_tensor::Tensor;
#[derive(Debug, Clone)]
pub struct LionConfig {
pub lr: f32,
pub beta1: f32,
pub beta2: f32,
pub weight_decay: f32,
}
impl Default for LionConfig {
fn default() -> Self {
Self {
lr: 1e-4,
beta1: 0.9,
beta2: 0.99,
weight_decay: 0.01,
}
}
}
pub struct Lion {
param_groups: Vec<ParamGroup>,
lr: f32,
beta1: f32,
beta2: f32,
weight_decay: f32,
momentum: HashMap<String, Tensor>,
step_count: usize,
}
impl Lion {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
beta1: f32,
beta2: f32,
weight_decay: f32,
) -> Self {
let param_group = ParamGroup::new(params, lr);
Self {
param_groups: vec![param_group],
lr,
beta1,
beta2,
weight_decay,
momentum: HashMap::new(),
step_count: 0,
}
}
pub fn from_config(params: Vec<Arc<RwLock<Tensor>>>, config: LionConfig) -> Self {
Self::new(
params,
config.lr,
config.beta1,
config.beta2,
config.weight_decay,
)
}
pub fn builder() -> LionBuilder {
LionBuilder::default()
}
pub fn set_lr_value(&mut self, lr: f32) {
self.lr = lr;
for group in &mut self.param_groups {
group.lr = lr;
}
}
pub fn set_beta1(&mut self, beta1: f32) {
self.beta1 = beta1;
}
pub fn set_beta2(&mut self, beta2: f32) {
self.beta2 = beta2;
}
pub fn set_weight_decay(&mut self, weight_decay: f32) {
self.weight_decay = weight_decay;
}
pub fn get_step_count(&self) -> usize {
self.step_count
}
}
impl Optimizer for Lion {
fn step(&mut self) -> OptimizerResult<()> {
self.step_count += 1;
for group in &self.param_groups {
let lr = group.lr;
let beta1 = self.beta1;
let beta2 = self.beta2;
let weight_decay = self.weight_decay;
for (idx, param) in group.params.iter().enumerate() {
let mut param_guard = param.write();
if !param_guard.has_grad() {
continue;
}
let grad = param_guard
.grad()
.ok_or_else(|| OptimizerError::InvalidInput("No gradient found".to_string()))?;
let param_key = format!("param_{}", idx);
let momentum_tensor = self.momentum.entry(param_key.clone()).or_insert_with(|| {
grad.zeros_like().expect("Failed to create momentum buffer")
});
let interpolated = momentum_tensor
.mul_scalar(beta1)
.map_err(|e| OptimizerError::TensorError(e))?
.add(
&grad
.mul_scalar(1.0 - beta1)
.map_err(|e| OptimizerError::TensorError(e))?,
)
.map_err(|e| OptimizerError::TensorError(e))?;
let update_direction = interpolated
.sign()
.map_err(|e| OptimizerError::TensorError(e))?;
let new_momentum = momentum_tensor
.mul_scalar(beta2)
.map_err(|e| OptimizerError::TensorError(e))?
.add(
&grad
.mul_scalar(1.0 - beta2)
.map_err(|e| OptimizerError::TensorError(e))?,
)
.map_err(|e| OptimizerError::TensorError(e))?;
*momentum_tensor = new_momentum;
let param_data = param_guard.clone();
let update = if weight_decay > 0.0 {
let decay_term = param_data
.mul_scalar(weight_decay)
.map_err(|e| OptimizerError::TensorError(e))?;
update_direction
.add(&decay_term)
.map_err(|e| OptimizerError::TensorError(e))?
.mul_scalar(lr)
.map_err(|e| OptimizerError::TensorError(e))?
} else {
update_direction
.mul_scalar(lr)
.map_err(|e| OptimizerError::TensorError(e))?
};
let new_param = param_data
.sub(&update)
.map_err(|e| OptimizerError::TensorError(e))?;
*param_guard = new_param;
}
}
Ok(())
}
fn zero_grad(&mut self) {
for group in &self.param_groups {
group.zero_grad();
}
}
fn get_lr(&self) -> Vec<f32> {
self.param_groups.iter().map(|g| g.lr).collect()
}
fn set_lr(&mut self, lr: f32) {
self.set_lr_value(lr);
}
fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
let lr = options.get("lr").copied().unwrap_or(self.lr);
let group = ParamGroup::new(params, lr).with_options(options);
self.param_groups.push(group);
}
fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
crate::optimizer::collect_parameters(&self.param_groups)
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let param_group_states = self
.param_groups
.iter()
.map(|g| ParamGroupState::from_param_group(g))
.collect();
let mut state = HashMap::new();
for (key, momentum) in &self.momentum {
let mut param_state = HashMap::new();
param_state.insert("momentum".to_string(), momentum.clone());
param_state.insert(
"step".to_string(),
Tensor::scalar(self.step_count as f32)
.map_err(|e| OptimizerError::TensorError(e))?,
);
state.insert(key.clone(), param_state);
}
let mut global_state = HashMap::new();
global_state.insert("beta1".to_string(), self.beta1);
global_state.insert("beta2".to_string(), self.beta2);
global_state.insert("weight_decay".to_string(), self.weight_decay);
Ok(OptimizerState {
optimizer_type: "Lion".to_string(),
version: "1.0".to_string(),
param_groups: param_group_states,
state,
global_state,
})
}
fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
if state.optimizer_type != "Lion" {
return Err(OptimizerError::InvalidInput(format!(
"Expected Lion state dict, got {}",
state.optimizer_type
)));
}
if let Some(&beta1) = state.global_state.get("beta1") {
self.beta1 = beta1;
}
if let Some(&beta2) = state.global_state.get("beta2") {
self.beta2 = beta2;
}
if let Some(&weight_decay) = state.global_state.get("weight_decay") {
self.weight_decay = weight_decay;
}
self.momentum.clear();
for (key, param_state) in state.state {
if let Some(momentum) = param_state.get("momentum") {
self.momentum.insert(key.clone(), momentum.clone());
}
if let Some(step_tensor) = param_state.get("step") {
self.step_count = step_tensor
.to_vec()
.map_err(|e| OptimizerError::TensorError(e))?[0]
as usize;
}
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct LionBuilder {
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
beta1: f32,
beta2: f32,
weight_decay: f32,
}
impl Default for LionBuilder {
fn default() -> Self {
let config = LionConfig::default();
Self {
params: Vec::new(),
lr: config.lr,
beta1: config.beta1,
beta2: config.beta2,
weight_decay: config.weight_decay,
}
}
}
impl LionBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn params(mut self, params: Vec<Arc<RwLock<Tensor>>>) -> Self {
self.params = params;
self
}
pub fn lr(mut self, lr: f32) -> Self {
self.lr = lr;
self
}
pub fn beta1(mut self, beta1: f32) -> Self {
self.beta1 = beta1;
self
}
pub fn beta2(mut self, beta2: f32) -> Self {
self.beta2 = beta2;
self
}
pub fn weight_decay(mut self, weight_decay: f32) -> Self {
self.weight_decay = weight_decay;
self
}
pub fn build(self) -> Lion {
Lion::new(
self.params,
self.lr,
self.beta1,
self.beta2,
self.weight_decay,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
use torsh_tensor::creation::randn;
#[test]
fn test_lion_creation() {
let param = Arc::new(RwLock::new(
randn::<f32>(&[2, 3]).expect("Failed to create tensor"),
));
let params = vec![param];
let optimizer = Lion::new(params, 1e-4, 0.9, 0.99, 0.01);
assert_eq!(optimizer.lr, 1e-4);
assert_eq!(optimizer.beta1, 0.9);
assert_eq!(optimizer.beta2, 0.99);
assert_eq!(optimizer.weight_decay, 0.01);
}
#[test]
fn test_lion_builder() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param];
let optimizer = Lion::builder()
.params(params)
.lr(2e-4)
.beta1(0.95)
.beta2(0.999)
.weight_decay(0.05)
.build();
assert_eq!(optimizer.lr, 2e-4);
assert_eq!(optimizer.beta1, 0.95);
assert_eq!(optimizer.beta2, 0.999);
assert_eq!(optimizer.weight_decay, 0.05);
Ok(())
}
#[test]
fn test_lion_step() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param.clone()];
let mut optimizer = Lion::new(params, 1e-4, 0.9, 0.99, 0.0);
let grad = randn::<f32>(&[2, 3])?;
param.write().set_grad(Some(grad.clone()));
let param_before = param.read().clone();
optimizer.step()?;
let param_after = param.read().clone();
let diff = param_before.sub(¶m_after)?;
let diff_norm = diff.norm()?.to_vec()?[0];
assert!(diff_norm > 0.0, "Parameters should have changed");
Ok(())
}
#[test]
fn test_lion_weight_decay() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param.clone()];
let mut optimizer = Lion::new(params, 1e-4, 0.9, 0.99, 0.1);
let grad = randn::<f32>(&[2, 3])?;
param.write().set_grad(Some(grad));
let param_before = param.read().clone();
optimizer.step()?;
let param_after = param.read().clone();
let param_norm_before = param_before.norm()?.to_vec()?[0];
let param_norm_after = param_after.norm()?.to_vec()?[0];
assert!(
(param_norm_before - param_norm_after).abs() > 0.0,
"Weight decay should affect parameters"
);
Ok(())
}
#[test]
fn test_lion_zero_grad() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param.clone()];
let mut optimizer = Lion::new(params, 1e-4, 0.9, 0.99, 0.01);
let grad = randn::<f32>(&[2, 3])?;
param.write().set_grad(Some(grad));
assert!(param.read().has_grad());
optimizer.zero_grad();
assert!(!param.read().has_grad());
Ok(())
}
#[test]
fn test_lion_state_dict() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param.clone()];
let mut optimizer = Lion::new(params, 1e-4, 0.9, 0.99, 0.01);
for _ in 0..3 {
let grad = randn::<f32>(&[2, 3])?;
param.write().set_grad(Some(grad));
optimizer.step()?;
optimizer.zero_grad();
}
let state = optimizer.state_dict()?;
assert_eq!(state.optimizer_type, "Lion");
assert_eq!(state.param_groups.len(), 1);
assert!(state.global_state.contains_key("beta1"));
assert!(state.global_state.contains_key("beta2"));
Ok(())
}
#[test]
fn test_lion_load_state_dict() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 3])?));
let params = vec![param.clone()];
let mut optimizer1 = Lion::new(params.clone(), 1e-4, 0.9, 0.99, 0.01);
for _ in 0..3 {
let grad = randn::<f32>(&[2, 3])?;
param.write().set_grad(Some(grad));
optimizer1.step()?;
optimizer1.zero_grad();
}
let state = optimizer1.state_dict()?;
let mut optimizer2 = Lion::new(params, 2e-4, 0.8, 0.98, 0.02);
optimizer2.load_state_dict(state)?;
assert_relative_eq!(optimizer2.beta1, 0.9, epsilon = 1e-6);
assert_relative_eq!(optimizer2.beta2, 0.99, epsilon = 1e-6);
Ok(())
}
}