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}