rlx_optim/qhadamw.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//! QHAdamW — Quasi-Hyperbolic Adam (Ma & Yarats, 2019) with decoupled
6//! weight decay.
7//!
8//! # Idea
9//!
10//! Adam can be viewed as "the second moment EMA scales the gradient,
11//! the first moment EMA *replaces* the gradient." The quasi-hyperbolic
12//! family says: don't *replace* — *interpolate*. Mix the EMA with the
13//! raw current gradient, controlled by per-moment scalars `ν₁, ν₂`.
14//!
15//! # Update rule
16//!
17//! ```text
18//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t
19//! v_t = β₂·v_{t-1} + (1 − β₂)·g_t²
20//! num = (1 − ν₁)·g_t + ν₁·m̂_t
21//! den = √((1 − ν₂)·g_t² + ν₂·v̂_t) + ε
22//! θ_t = θ_{t-1} − lr · ( num / den + λ·θ_{t-1} )
23//! ```
24//!
25//! Setting `ν₁ = ν₂ = 1` recovers standard AdamW; `ν₁ = β₁` recovers
26//! Nesterov-style behavior in the limit; `ν₁ < 1` injects more of the
27//! current gradient and tends to be more robust on noisy losses.
28//!
29//! # When to use
30//!
31//! When you've found AdamW too sluggish on noisy / heavy-tail
32//! gradients (RL, GAN training) and an LR sweep didn't help.
33
34use std::collections::HashMap;
35
36use crate::Optimizer;
37use crate::common::zeros_entry;
38
39/// Quasi-hyperbolic AdamW. Per-tensor state: two `f32` buffers.
40#[derive(Debug, Clone)]
41pub struct QHAdamW {
42 /// Learning rate.
43 pub lr: f32,
44 /// First-moment EMA decay β₁. Ma & Yarats recommend `0.995` (much
45 /// closer to 1 than vanilla Adam) — the QH interpolation already
46 /// keeps current-gradient weight in the numerator.
47 pub beta1: f32,
48 /// Second-moment EMA decay β₂. Default `0.999`.
49 pub beta2: f32,
50 /// First-moment QH interpolation coefficient ν₁ ∈ \[0, 1\].
51 /// `1.0` = pure EMA (standard Adam first moment); `0.0` = pure
52 /// current gradient (no momentum). Default `0.7`.
53 pub nu1: f32,
54 /// Second-moment QH interpolation coefficient ν₂. `1.0` = standard
55 /// Adam denominator. Default `1.0`.
56 pub nu2: f32,
57 /// Denominator stability constant. Default `1e-8`.
58 pub eps: f32,
59 /// Decoupled weight-decay coefficient λ. Default `0.01`.
60 pub weight_decay: f32,
61 step: u64,
62 m: HashMap<String, Vec<f32>>,
63 v: HashMap<String, Vec<f32>>,
64}
65
66impl QHAdamW {
67 /// Construct with `(β₁, β₂, ν₁, ν₂, ε, λ) = (0.995, 0.999, 0.7, 1.0, 1e-8, 0.01)`.
68 pub fn new(lr: f32) -> Self {
69 Self {
70 lr,
71 beta1: 0.995,
72 beta2: 0.999,
73 nu1: 0.7,
74 nu2: 1.0,
75 eps: 1e-8,
76 weight_decay: 0.01,
77 step: 0,
78 m: HashMap::new(),
79 v: HashMap::new(),
80 }
81 }
82
83 /// Override (β₁, β₂).
84 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
85 self.beta1 = b1;
86 self.beta2 = b2;
87 self
88 }
89
90 /// Override the quasi-hyperbolic coefficients (ν₁, ν₂).
91 pub fn with_nus(mut self, n1: f32, n2: f32) -> Self {
92 self.nu1 = n1;
93 self.nu2 = n2;
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 QHAdamW {
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 n1 = self.nu1 as f64;
117 let n2 = self.nu2 as f64;
118 let eps = self.eps as f64;
119 let lr = self.lr as f64;
120 let wd = self.weight_decay as f64;
121 let m = zeros_entry(&mut self.m, name, param.len());
122 let v = zeros_entry(&mut self.v, name, param.len());
123 for i in 0..param.len() {
124 let g = grad[i] as f64;
125 let mi = b1 * m[i] as f64 + (1.0 - b1) * g;
126 let vi = b2 * v[i] as f64 + (1.0 - b2) * g * g;
127 m[i] = mi as f32;
128 v[i] = vi as f32;
129 let m_hat = mi / bc1;
130 let v_hat = vi / bc2;
131 // Quasi-hyperbolic numerator & denominator (Ma & Yarats Alg. 2).
132 let num = (1.0 - n1) * g + n1 * m_hat;
133 let den = ((1.0 - n2) * g * g + n2 * v_hat).sqrt() + eps;
134 let p = param[i] as f64;
135 param[i] = (p - lr * (num / den + wd * p)) as f32;
136 }
137 }
138
139 fn end_iteration(&mut self) {
140 self.step += 1;
141 }
142}