rlx_optim/mars.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//! MARS — Make vAriance Reduction Shine (Yuan, Liu, Wu, Su, Gu, 2024).
6//!
7//! # Idea
8//!
9//! Variance reduction (SVRG, SARAH) lowers gradient noise by mixing in
10//! a *previous* gradient — at the cost of an extra forward/backward
11//! pass per snapshot. MARS shows that you don't need a snapshot:
12//! using just `g_{t−1}` (the previous mini-batch's gradient) as the
13//! "control variate" gives most of the benefit of full variance
14//! reduction, for free.
15//!
16//! # Update rule
17//!
18//! ```text
19//! c_t = g_t + γ · β₁/(1−β₁) · (g_t − g_{t−1}) // VR-corrected grad
20//! [optional: clip c_t to unit norm per tensor]
21//! m_t = β₁·m_{t-1} + (1 − β₁)·c_t
22//! v_t = β₂·v_{t-1} + (1 − β₂)·c_t²
23//! θ_t = θ_{t-1} − lr · ( m̂_t/(√v̂_t + ε) + λ·θ_{t-1} ) // AdamW-style
24//! ```
25//!
26//! γ = 0 collapses MARS to AdamW. γ = 1 is the "full" Yuan et al.
27//! prescription; the recommended sweet spot is `γ ≈ 0.025`.
28//!
29//! # When to use
30//!
31//! Anywhere variance-reduced SGD/Adam variants would help — noisy
32//! gradients, small batches, RL-style on-policy training. State cost
33//! per parameter: three buffers (`m`, `v`, previous-gradient cache).
34
35use std::collections::HashMap;
36
37use crate::Optimizer;
38use crate::common::zeros_entry;
39
40/// MARS — variance-reduced AdamW. Per-tensor state: three `f32`
41/// buffers (`m`, `v`, previous-gradient cache).
42#[derive(Debug, Clone)]
43pub struct Mars {
44 /// Learning rate.
45 pub lr: f32,
46 /// First-moment EMA decay β₁. Default `0.95`.
47 pub beta1: f32,
48 /// Second-moment EMA decay β₂. Default `0.99`.
49 pub beta2: f32,
50 /// Denominator stability constant. Default `1e-8`.
51 pub eps: f32,
52 /// Decoupled weight-decay coefficient λ. Default `0.0`.
53 pub weight_decay: f32,
54 /// Variance-reduction strength γ (Yuan et al. eq. 7). `0.0`
55 /// collapses MARS to plain AdamW; `1.0` is the full prescription.
56 /// Default `0.025`.
57 pub gamma: f32,
58 /// If `true`, clip the variance-reduced surrogate `c_t` to unit
59 /// norm per tensor (matches the "MARS-AdamW" recipe and keeps
60 /// the VR kick from exploding when `g_{t-1}` is unrelated noise
61 /// on early steps).
62 pub clip_c: bool,
63 step: u64,
64 m: HashMap<String, Vec<f32>>,
65 v: HashMap<String, Vec<f32>>,
66 prev_g: HashMap<String, Vec<f32>>,
67 /// Reusable scratch for the variance-reduced surrogate `c_t`.
68 scratch: HashMap<String, Vec<f32>>,
69}
70
71impl Mars {
72 /// Construct with `(β₁, β₂, ε, λ, γ, clip_c) = (0.95, 0.99, 1e-8, 0.0, 0.025, true)`.
73 pub fn new(lr: f32) -> Self {
74 Self {
75 lr,
76 beta1: 0.95,
77 beta2: 0.99,
78 eps: 1e-8,
79 weight_decay: 0.0,
80 gamma: 0.025,
81 clip_c: true,
82 step: 0,
83 m: HashMap::new(),
84 v: HashMap::new(),
85 prev_g: HashMap::new(),
86 scratch: HashMap::new(),
87 }
88 }
89
90 /// Override (β₁, β₂).
91 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
92 self.beta1 = b1;
93 self.beta2 = b2;
94 self
95 }
96
97 /// Override the decoupled-decay coefficient.
98 pub fn with_weight_decay(mut self, wd: f32) -> Self {
99 self.weight_decay = wd;
100 self
101 }
102}
103
104impl Optimizer for Mars {
105 fn set_lr(&mut self, lr: f32) {
106 self.lr = lr;
107 }
108
109 fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
110 debug_assert_eq!(param.len(), grad.len());
111 let t = (self.step + 1) as f64;
112 let b1 = self.beta1 as f64;
113 let b2 = self.beta2 as f64;
114 let bc1 = 1.0 - b1.powf(t);
115 let bc2 = 1.0 - b2.powf(t);
116 let scale = self.gamma as f64 * b1 / (1.0 - b1);
117 let eps = self.eps as f64;
118 let lr = self.lr as f64;
119 let wd = self.weight_decay as f64;
120 // Four distinct fields ⇒ borrows can coexist.
121 let prev = zeros_entry(&mut self.prev_g, name, param.len());
122 let c = zeros_entry(&mut self.scratch, name, param.len());
123 let m = zeros_entry(&mut self.m, name, param.len());
124 let v = zeros_entry(&mut self.v, name, param.len());
125 // c_t = g_t + scale * (g_t - g_{t-1})
126 let mut c_sq_norm = 0.0f64;
127 for i in 0..param.len() {
128 let g = grad[i] as f64;
129 let pg = prev[i] as f64;
130 let ci = g + scale * (g - pg);
131 c[i] = ci as f32;
132 c_sq_norm += ci * ci;
133 prev[i] = grad[i];
134 }
135 // Optional per-tensor norm clip on c (keeps the VR kick from
136 // exploding on early steps when g_{t-1} is unrelated noise).
137 if self.clip_c && c_sq_norm > 1.0 {
138 let s = (1.0 / c_sq_norm.sqrt()) as f32;
139 for ci in c.iter_mut() {
140 *ci *= s;
141 }
142 }
143 for i in 0..param.len() {
144 let ci = c[i] as f64;
145 let mi = b1 * m[i] as f64 + (1.0 - b1) * ci;
146 let vi = b2 * v[i] as f64 + (1.0 - b2) * ci * ci;
147 m[i] = mi as f32;
148 v[i] = vi as f32;
149 let m_hat = mi / bc1;
150 let v_hat = vi / bc2;
151 let p = param[i] as f64;
152 param[i] = (p - lr * (m_hat / (v_hat.sqrt() + eps) + wd * p)) as f32;
153 }
154 }
155
156 fn end_iteration(&mut self) {
157 self.step += 1;
158 }
159}