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}