1use 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#[derive(Debug, Clone)]
41pub struct Lion<A: Float + ScalarOperand + Debug> {
42 learning_rate: A,
44 beta1: A,
46 beta2: A,
48 weight_decay: A,
50 m: Option<Vec<Array<A, IxDyn>>>,
52}
53
54impl<A: Float + ScalarOperand + Debug + Send + Sync> Lion<A> {
55 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 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 pub fn set_beta1(&mut self, beta1: A) -> &mut Self {
90 self.beta1 = beta1;
91 self
92 }
93
94 pub fn get_beta1(&self) -> A {
96 self.beta1
97 }
98
99 pub fn set_beta2(&mut self, beta2: A) -> &mut Self {
101 self.beta2 = beta2;
102 self
103 }
104
105 pub fn get_beta2(&self) -> A {
107 self.beta2
108 }
109
110 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
112 self.weight_decay = weight_decay;
113 self
114 }
115
116 pub fn get_weight_decay(&self) -> A {
118 self.weight_decay
119 }
120
121 pub fn learning_rate(&self) -> A {
123 self.learning_rate
124 }
125
126 pub fn set_lr(&mut self, lr: A) {
128 self.learning_rate = lr;
129 }
130
131 pub fn reset(&mut self) {
133 self.m = None;
134 }
135
136 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 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 let interpolated = *m_i * beta1 + g * (one - beta1);
188
189 let sign_update = if interpolated > zero {
191 one
192 } else if interpolated < zero {
193 -one
194 } else {
195 zero
196 };
197
198 let decayed = if use_weight_decay {
200 *p * decay_factor
201 } else {
202 *p
203 };
204 *p = decayed - sign_update * lr;
205
206 *m_i = *m_i * beta2 + g * (one - beta2);
208 });
209
210 Ok(())
211 }
212
213 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 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); let mut params = Array1::from_vec(vec![5.0]);
296
297 for _ in 0..40 {
299 let gradients = Array1::from_vec(vec![2.0 * params[0]]);
302 params = optimizer
303 .step(¶ms, &gradients)
304 .expect("optimizer.step succeeds in test_lion_convergence");
305 }
306
307 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 let params = Array1::from_vec(vec![1.0]);
317 let gradients = Array1::from_vec(vec![0.1]);
318 let _ = optimizer
319 .step(¶ms, &gradients)
320 .expect("optimizer.step succeeds in test_lion_reset");
321
322 optimizer.reset();
324
325 let next_step = optimizer
327 .step(¶ms, &gradients)
328 .expect("optimizer.step succeeds in test_lion_reset");
329
330 let mut fresh_optimizer: Lion<f64> = Lion::new(0.1);
332 let fresh_step = fresh_optimizer
333 .step(¶ms, &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}