Skip to main content

optirs_core/optimizers/
lion.rs

1// Lion optimizer implementation
2//
3// Based on the paper "Symbolic Discovery of Optimization Algorithms"
4// by Chen et al. (2023).
5
6use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
7use scirs2_core::numeric::Float;
8use std::fmt::Debug;
9
10use crate::error::{OptimError, Result};
11use crate::optimizers::Optimizer;
12
13/// Lion optimizer
14///
15/// Implements the Lion (Evolved Sign Momentum) optimization algorithm.
16/// Lion is a memory-efficient optimizer that achieves strong performance
17/// with only momentum state and uses the sign of the momentum for updates.
18///
19/// Formula:
20/// u_t = beta1 * m_{t-1} + (1 - beta1) * g_t
21/// theta_t = theta_{t-1} - alpha * (sign(u_t) + lambda * theta_{t-1})
22/// m_t = beta2 * m_{t-1} + (1 - beta2) * g_t
23///
24/// # Examples
25///
26/// ```
27/// use scirs2_core::ndarray::Array1;
28/// use optirs_core::optimizers::{Lion, Optimizer};
29///
30/// // Initialize parameters and gradients
31/// let params = Array1::zeros(5);
32/// let gradients = Array1::from_vec(vec![0.1, 0.2, -0.3, 0.0, 0.5]);
33///
34/// // Create a Lion optimizer with default hyperparameters
35/// let mut optimizer = Lion::new(0.001);
36///
37/// // Update parameters
38/// let new_params = optimizer.step(&params, &gradients).expect("optimizer.step succeeds");
39/// ```
40#[derive(Debug, Clone)]
41pub struct Lion<A: Float + ScalarOperand + Debug> {
42    /// Learning rate
43    learning_rate: A,
44    /// Exponential decay rate for the momentum
45    beta1: A,
46    /// Exponential decay rate for the momentum update
47    beta2: A,
48    /// Weight decay factor (L2 regularization)
49    weight_decay: A,
50    /// Momentum vectors, one slot per parameter-tensor index
51    m: Option<Vec<Array<A, IxDyn>>>,
52}
53
54impl<A: Float + ScalarOperand + Debug + Send + Sync> Lion<A> {
55    /// Creates a new Lion optimizer with the given learning rate and default settings
56    ///
57    /// # Arguments
58    ///
59    /// * `learning_rate` - The learning rate for parameter updates
60    pub fn new(learning_rate: A) -> Self {
61        Self {
62            learning_rate,
63            beta1: A::from(0.9).expect("Lion: default beta1 (0.9) must fit in A"),
64            beta2: A::from(0.99).expect("Lion: default beta2 (0.99) must fit in A"),
65            weight_decay: A::zero(),
66            m: None,
67        }
68    }
69
70    /// Creates a new Lion optimizer with the full configuration
71    ///
72    /// # Arguments
73    ///
74    /// * `learning_rate` - The learning rate for parameter updates
75    /// * `beta1` - Exponential decay rate for computing the interpolated update (default: 0.9)
76    /// * `beta2` - Exponential decay rate for updating the momentum (default: 0.99)
77    /// * `weight_decay` - Weight decay factor for L2 regularization (default: 0.0)
78    pub fn new_with_config(learning_rate: A, beta1: A, beta2: A, weight_decay: A) -> Self {
79        Self {
80            learning_rate,
81            beta1,
82            beta2,
83            weight_decay,
84            m: None,
85        }
86    }
87
88    /// Sets the beta1 parameter
89    pub fn set_beta1(&mut self, beta1: A) -> &mut Self {
90        self.beta1 = beta1;
91        self
92    }
93
94    /// Gets the beta1 parameter
95    pub fn get_beta1(&self) -> A {
96        self.beta1
97    }
98
99    /// Sets the beta2 parameter
100    pub fn set_beta2(&mut self, beta2: A) -> &mut Self {
101        self.beta2 = beta2;
102        self
103    }
104
105    /// Gets the beta2 parameter
106    pub fn get_beta2(&self) -> A {
107        self.beta2
108    }
109
110    /// Sets the weight decay parameter
111    pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
112        self.weight_decay = weight_decay;
113        self
114    }
115
116    /// Gets the weight decay parameter
117    pub fn get_weight_decay(&self) -> A {
118        self.weight_decay
119    }
120
121    /// Gets the current learning rate
122    pub fn learning_rate(&self) -> A {
123        self.learning_rate
124    }
125
126    /// Sets the learning rate
127    pub fn set_lr(&mut self, lr: A) {
128        self.learning_rate = lr;
129    }
130
131    /// Resets the internal state of the optimizer
132    pub fn reset(&mut self) {
133        self.m = None;
134    }
135
136    /// Ensures a momentum slot exists for `index` and matches `dim`
137    fn ensure_state(&mut self, index: usize, dim: &IxDyn) {
138        let m = self.m.get_or_insert_with(Vec::new);
139        while m.len() <= index {
140            m.push(Array::zeros(dim.clone()));
141        }
142        if m[index].raw_dim() != *dim {
143            m[index] = Array::zeros(dim.clone());
144        }
145    }
146
147    /// Applies a Lion update in place for the parameter tensor at `index`
148    pub fn step_inplace_indexed<D: Dimension>(
149        &mut self,
150        index: usize,
151        params: &mut Array<A, D>,
152        gradients: &Array<A, D>,
153    ) -> Result<()> {
154        if params.shape() != gradients.shape() {
155            return Err(OptimError::DimensionMismatch(format!(
156                "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
157                params.shape(),
158                gradients.shape()
159            )));
160        }
161
162        let dim = params.raw_dim().into_dyn();
163        self.ensure_state(index, &dim);
164
165        let beta1 = self.beta1;
166        let beta2 = self.beta2;
167        let lr = self.learning_rate;
168        let weight_decay = self.weight_decay;
169        let use_weight_decay = weight_decay > A::zero();
170        let one = A::one();
171        let zero = A::zero();
172        let decay_factor = one - weight_decay * lr;
173
174        let m = self
175            .m
176            .as_mut()
177            .ok_or_else(|| OptimError::InvalidConfig("Lion state not initialized".to_string()))?;
178
179        let mut params_view = params.view_mut().into_dyn();
180        let gradients_view = gradients.view().into_dyn();
181
182        Zip::from(&mut params_view)
183            .and(&gradients_view)
184            .and(&mut m[index])
185            .for_each(|p, &g, m_i| {
186                // Step 1: interpolated update using beta1
187                let interpolated = *m_i * beta1 + g * (one - beta1);
188
189                // Step 2: sign of the interpolated update
190                let sign_update = if interpolated > zero {
191                    one
192                } else if interpolated < zero {
193                    -one
194                } else {
195                    zero
196                };
197
198                // Step 3: decoupled weight decay, then the sign step
199                let decayed = if use_weight_decay {
200                    *p * decay_factor
201                } else {
202                    *p
203                };
204                *p = decayed - sign_update * lr;
205
206                // Step 4: momentum update using beta2
207                *m_i = *m_i * beta2 + g * (one - beta2);
208            });
209
210        Ok(())
211    }
212
213    /// Applies a Lion update in place using the state slot of the first parameter tensor
214    pub fn step_inplace<D: Dimension>(
215        &mut self,
216        params: &mut Array<A, D>,
217        gradients: &Array<A, D>,
218    ) -> Result<()> {
219        self.step_inplace_indexed(0, params, gradients)
220    }
221
222    /// Performs a Lion update for the parameter tensor at `index`
223    ///
224    /// Each `index` owns an independent momentum slot.
225    pub fn step_indexed<D: Dimension>(
226        &mut self,
227        index: usize,
228        params: &Array<A, D>,
229        gradients: &Array<A, D>,
230    ) -> Result<Array<A, D>> {
231        let mut updated = params.to_owned();
232        self.step_inplace_indexed(index, &mut updated, gradients)?;
233        Ok(updated)
234    }
235}
236
237impl<A, D> Optimizer<A, D> for Lion<A>
238where
239    A: Float + ScalarOperand + Debug + Send + Sync,
240    D: Dimension,
241{
242    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
243        self.step_indexed(0, params, gradients)
244    }
245
246    fn step_list(
247        &mut self,
248        params_list: &[&Array<A, D>],
249        gradients_list: &[&Array<A, D>],
250    ) -> Result<Vec<Array<A, D>>> {
251        if params_list.len() != gradients_list.len() {
252            return Err(OptimError::InvalidConfig(format!(
253                "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
254                params_list.len(),
255                gradients_list.len()
256            )));
257        }
258
259        let mut results = Vec::with_capacity(params_list.len());
260        for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
261            results.push(self.step_indexed(index, params, grads)?);
262        }
263        Ok(results)
264    }
265
266    fn get_learning_rate(&self) -> A {
267        self.learning_rate
268    }
269
270    fn set_learning_rate(&mut self, learning_rate: A) {
271        self.learning_rate = learning_rate;
272    }
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278    use approx::assert_abs_diff_eq;
279    use scirs2_core::ndarray::Array1;
280
281    #[test]
282    fn test_lion_basic_creation() {
283        let optimizer: Lion<f64> = Lion::new(0.001);
284        assert_abs_diff_eq!(optimizer.learning_rate(), 0.001);
285        assert_abs_diff_eq!(optimizer.get_beta1(), 0.9);
286        assert_abs_diff_eq!(optimizer.get_beta2(), 0.99);
287        assert_abs_diff_eq!(optimizer.get_weight_decay(), 0.0);
288    }
289
290    #[test]
291    fn test_lion_convergence() {
292        let mut optimizer: Lion<f64> = Lion::new(0.1); // Higher learning rate for testing
293
294        // Minimize a simple quadratic function: f(x) = x^2
295        let mut params = Array1::from_vec(vec![5.0]);
296
297        // Lion converges linearly with sign updates
298        for _ in 0..40 {
299            // Fewer iterations with higher learning rate
300            // Gradient of x^2 is 2x
301            let gradients = Array1::from_vec(vec![2.0 * params[0]]);
302            params = optimizer
303                .step(&params, &gradients)
304                .expect("optimizer.step succeeds in test_lion_convergence");
305        }
306
307        // With learning rate 0.1 and 40 iterations, should reach close to 1.0
308        assert!(params[0].abs() < 1.1);
309    }
310
311    #[test]
312    fn test_lion_reset() {
313        let mut optimizer: Lion<f64> = Lion::new(0.1);
314
315        // Perform a step to initialize state
316        let params = Array1::from_vec(vec![1.0]);
317        let gradients = Array1::from_vec(vec![0.1]);
318        let _ = optimizer
319            .step(&params, &gradients)
320            .expect("optimizer.step succeeds in test_lion_reset");
321
322        // Reset optimizer
323        optimizer.reset();
324
325        // Next step should behave like the first
326        let next_step = optimizer
327            .step(&params, &gradients)
328            .expect("optimizer.step succeeds in test_lion_reset");
329
330        // Create fresh optimizer for comparison
331        let mut fresh_optimizer: Lion<f64> = Lion::new(0.1);
332        let fresh_step = fresh_optimizer
333            .step(&params, &gradients)
334            .expect("step succeeds in test_lion_reset");
335
336        assert_abs_diff_eq!(next_step[0], fresh_step[0], epsilon = 1e-10);
337    }
338}