Skip to main content

rlx_optim/
sgd.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//! Stochastic Gradient Descent with optional momentum and decoupled
6//! L2 weight decay.
7//!
8//! # Update rules
9//!
10//! Vanilla SGD (`momentum = 0`):
11//!
12//! ```text
13//! θ_{t+1} = θ_t − lr · (g_t + λ·θ_t)
14//! ```
15//!
16//! Polyak momentum (`momentum = μ`, `nesterov = false`):
17//!
18//! ```text
19//! v_{t+1} = μ·v_t + (g_t + λ·θ_t)
20//! θ_{t+1} = θ_t − lr · v_{t+1}
21//! ```
22//!
23//! Nesterov-accelerated SGD (`nesterov = true`):
24//!
25//! ```text
26//! v_{t+1} = μ·v_t + (g_t + λ·θ_t)
27//! θ_{t+1} = θ_t − lr · (g_t + λ·θ_t + μ·v_{t+1})
28//! ```
29//!
30//! # When to use
31//!
32//! The default choice when training CNNs from scratch; with a
33//! well-tuned `lr` schedule it still beats Adam on many vision
34//! benchmarks. Cheap state (one buffer if `momentum > 0`).
35
36use std::collections::HashMap;
37
38use crate::Optimizer;
39use crate::common::zeros_entry;
40
41/// SGD with momentum / Nesterov / L2 weight decay.
42///
43/// All hyperparameters are public so callers can hot-swap them between
44/// iterations (e.g. for a warm-up schedule). State is keyed by
45/// parameter name; the same `Sgd` instance can drive every tensor in
46/// a model.
47#[derive(Debug, Clone)]
48pub struct Sgd {
49    /// Learning rate. No default — pass it to [`Sgd::new`].
50    pub lr: f32,
51    /// Polyak momentum coefficient ∈ \[0, 1\). `0.0` disables momentum
52    /// entirely (and the per-tensor velocity buffer is still allocated
53    /// but unused — set via [`Sgd::with_momentum`] if you want it on).
54    pub momentum: f32,
55    /// Use Nesterov-accelerated momentum. Only meaningful when
56    /// `momentum > 0`.
57    pub nesterov: bool,
58    /// L2 weight decay coefficient λ. Folded into the gradient
59    /// *before* the momentum EMA (classical, **not** decoupled).
60    /// Use [`crate::AdamW`]-style decoupling if you need that.
61    pub weight_decay: f32,
62    v: HashMap<String, Vec<f32>>,
63}
64
65impl Sgd {
66    /// Construct with `lr` and momentum / decay disabled.
67    pub fn new(lr: f32) -> Self {
68        Self {
69            lr,
70            momentum: 0.0,
71            nesterov: false,
72            weight_decay: 0.0,
73            v: HashMap::new(),
74        }
75    }
76
77    /// Enable Polyak (or Nesterov) momentum.
78    pub fn with_momentum(mut self, momentum: f32, nesterov: bool) -> Self {
79        self.momentum = momentum;
80        self.nesterov = nesterov;
81        self
82    }
83
84    /// Set the L2 weight-decay coefficient.
85    pub fn with_weight_decay(mut self, wd: f32) -> Self {
86        self.weight_decay = wd;
87        self
88    }
89}
90
91// ── checkpointing ───────────────────────────────────────────
92
93impl Sgd {
94    /// Named accumulators plus the step counter — see
95    /// [`crate::OptimizerState`].
96    pub(crate) fn snapshot(&self) -> crate::OptimizerState {
97        let mut out = crate::OptimizerState {
98            step: 0,
99            buffers: Vec::new(),
100        };
101        out.extend_slot("v", &self.v);
102        out
103    }
104
105    pub(crate) fn restore(&mut self, state: &crate::OptimizerState) {
106        state.take_slot("v", &mut self.v);
107    }
108}
109
110impl Optimizer for Sgd {
111    fn set_lr(&mut self, lr: f32) {
112        self.lr = lr;
113    }
114
115    fn state_dict(&self) -> Option<crate::OptimizerState> {
116        Some(self.snapshot())
117    }
118
119    fn load_state_dict(&mut self, state: &crate::OptimizerState) -> bool {
120        self.restore(state);
121        true
122    }
123
124    fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
125        debug_assert_eq!(param.len(), grad.len());
126        let mu = self.momentum;
127        let wd = self.weight_decay;
128        let lr = self.lr;
129        // Read out of `self` *before* the loop. Both of these used to be tested
130        // per element, and `self.nesterov` was read through a live `&mut self`
131        // borrow, which stopped the loop vectorizing: SGD measured 1.06
132        // ns/element against Lion's 0.28 while doing strictly less arithmetic.
133        let nesterov = self.nesterov;
134        let v = zeros_entry(&mut self.v, name, param.len());
135
136        // One specialized loop per configuration rather than one loop with the
137        // configuration inside it. The arithmetic is unchanged, so results stay
138        // bit-identical.
139        if mu == 0.0 {
140            for (p, g) in param.iter_mut().zip(grad) {
141                let d = *g + wd * *p;
142                *p -= lr * d;
143            }
144        } else if nesterov {
145            for ((p, vi), g) in param.iter_mut().zip(v.iter_mut()).zip(grad) {
146                let d = *g + wd * *p;
147                *vi = mu * *vi + d;
148                *p -= lr * (d + mu * *vi);
149            }
150        } else {
151            for ((p, vi), g) in param.iter_mut().zip(v.iter_mut()).zip(grad) {
152                let d = *g + wd * *p;
153                *vi = mu * *vi + d;
154                *p -= lr * *vi;
155            }
156        }
157    }
158}