Skip to main content

rlx_optim/
lion.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3// SPDX-License-Identifier: MIT OR Apache-2.0
4
5//! Lion — EvoLved Sign Momentum (Chen et al., 2023, "Symbolic
6//! Discovery of Optimization Algorithms").
7//!
8//! # Idea
9//!
10//! Lion was *discovered* by a program-synthesis search over candidate
11//! optimizer expressions. The found rule is shockingly simple — one
12//! momentum buffer, and the update is the **sign** of an
13//! interpolation between the momentum and the gradient.
14//!
15//! # Update rule
16//!
17//! ```text
18//! c_t   = β₁·m_{t-1} + (1 − β₁)·g_t
19//! θ_t   = θ_{t-1} − lr · ( sign(c_t) + λ·θ_{t-1} )
20//! m_t   = β₂·m_{t-1} + (1 − β₂)·g_t          // note: different β₂!
21//! ```
22//!
23//! Two distinct betas: `β₁` shapes the *update direction* (faster
24//! adaptation), `β₂` shapes the *carried momentum* (slower memory).
25//!
26//! # When to use
27//!
28//! Half the memory of Adam (one buffer instead of two), often
29//! converges to similar quality on transformers when the LR is
30//! tuned 3–10× lower than the corresponding AdamW LR. Sign updates
31//! get coarse on tiny problems — favor large-batch / large-model
32//! regimes.
33
34use std::collections::HashMap;
35
36use crate::Optimizer;
37use crate::common::{zeros_entry, zip3_for_each};
38
39/// EvoLved sign-momentum optimizer.
40///
41/// Per-tensor state: **one** `f32` buffer (half of Adam's footprint).
42#[derive(Debug, Clone)]
43pub struct Lion {
44    /// Learning rate. **Critical**: typically 3–10× smaller than the
45    /// AdamW LR you'd use on the same model (because the update has
46    /// unit `‖sign(·)‖` per coordinate).
47    pub lr: f32,
48    /// Interpolation coefficient for the *update direction* (β₁ in
49    /// Chen et al.). Default `0.9`.
50    pub beta1: f32,
51    /// EMA coefficient for the *carried momentum* (β₂). Default `0.99`.
52    pub beta2: f32,
53    /// Decoupled weight-decay coefficient λ. Tune ~3–10× higher than
54    /// the AdamW λ you'd pair with the same model. Default `0.0`.
55    pub weight_decay: f32,
56    m: HashMap<String, Vec<f32>>,
57}
58
59impl Lion {
60    /// Construct with `(β₁, β₂, λ) = (0.9, 0.99, 0.0)`.
61    pub fn new(lr: f32) -> Self {
62        Self {
63            lr,
64            beta1: 0.9,
65            beta2: 0.99,
66            weight_decay: 0.0,
67            m: HashMap::new(),
68        }
69    }
70
71    /// Override (β₁, β₂). They serve different roles — see the
72    /// struct-level docs.
73    pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
74        self.beta1 = b1;
75        self.beta2 = b2;
76        self
77    }
78
79    /// Override the decoupled-decay coefficient.
80    pub fn with_weight_decay(mut self, wd: f32) -> Self {
81        self.weight_decay = wd;
82        self
83    }
84}
85
86// ── checkpointing ───────────────────────────────────────────
87
88impl Lion {
89    /// Named accumulators plus the step counter — see
90    /// [`crate::OptimizerState`].
91    pub(crate) fn snapshot(&self) -> crate::OptimizerState {
92        let mut out = crate::OptimizerState {
93            step: 0,
94            buffers: Vec::new(),
95        };
96        out.extend_slot("m", &self.m);
97        out
98    }
99
100    pub(crate) fn restore(&mut self, state: &crate::OptimizerState) {
101        state.take_slot("m", &mut self.m);
102    }
103}
104
105impl Optimizer for Lion {
106    fn set_lr(&mut self, lr: f32) {
107        self.lr = lr;
108    }
109
110    fn state_dict(&self) -> Option<crate::OptimizerState> {
111        Some(self.snapshot())
112    }
113
114    fn load_state_dict(&mut self, state: &crate::OptimizerState) -> bool {
115        self.restore(state);
116        true
117    }
118
119    fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
120        debug_assert_eq!(param.len(), grad.len());
121        let b1 = self.beta1;
122        let b2 = self.beta2;
123        let lr = self.lr;
124        let wd = self.weight_decay;
125        let m = zeros_entry(&mut self.m, name, param.len());
126        zip3_for_each(param, m, grad, |p, mi, gi| {
127            // Update direction = sign(b1*m + (1-b1)*g)
128            let c = b1 * *mi + (1.0 - b1) * gi;
129            let sign = if c > 0.0 {
130                1.0
131            } else if c < 0.0 {
132                -1.0
133            } else {
134                0.0
135            };
136            // Decoupled weight decay (matches Chen et al. eq. 1).
137            *p -= lr * (sign + wd * *p);
138            // Then update the momentum with a different β₂.
139            *mi = b2 * *mi + (1.0 - b2) * gi;
140        });
141    }
142}