rlx_optim/nadamw.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//! NAdamW — Nesterov-accelerated AdamW.
6//!
7//! Combines Dozat (2016, "Incorporating Nesterov Momentum into Adam")
8//! with the decoupled-decay rule of [`crate::AdamW`]. The Nesterov
9//! trick is a *lookahead*: instead of using the bias-corrected first
10//! moment `m̂_t` directly, we use a convex combination of `m̂_{t+1}`
11//! (predicted) and the current bias-corrected gradient.
12//!
13//! # Update rule
14//!
15//! ```text
16//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t
17//! v_t = β₂·v_{t-1} + (1 − β₂)·g_t²
18//! m̄_t = β₁ · m_t / (1 − β₁^{t+1}) + (1 − β₁) · g_t / (1 − β₁ᵗ)
19//! v̂_t = v_t / (1 − β₂ᵗ)
20//! θ_t = θ_{t-1} − lr · ( m̄_t/(√v̂_t + ε) + λ·θ_{t-1} )
21//! ```
22//!
23//! # When to use
24//!
25//! Slightly more aggressive than AdamW on the early steps; the
26//! lookahead first-moment occasionally helps escape flat regions.
27//! Same state cost as AdamW.
28
29use std::collections::HashMap;
30
31use crate::Optimizer;
32use crate::common::zeros_entry;
33
34/// Nesterov AdamW. Per-tensor state: two `f32` buffers.
35#[derive(Debug, Clone)]
36pub struct NAdamW {
37 /// Learning rate.
38 pub lr: f32,
39 /// First-moment EMA decay β₁. Default `0.9`.
40 pub beta1: f32,
41 /// Second-moment EMA decay β₂. Default `0.999`.
42 pub beta2: f32,
43 /// Denominator stability constant. Default `1e-8`.
44 pub eps: f32,
45 /// Decoupled weight-decay coefficient λ. Default `0.01`.
46 pub weight_decay: f32,
47 step: u64,
48 m: HashMap<String, Vec<f32>>,
49 v: HashMap<String, Vec<f32>>,
50}
51
52impl NAdamW {
53 /// Construct with `(β₁, β₂, ε, λ) = (0.9, 0.999, 1e-8, 0.01)`.
54 pub fn new(lr: f32) -> Self {
55 Self {
56 lr,
57 beta1: 0.9,
58 beta2: 0.999,
59 eps: 1e-8,
60 weight_decay: 0.01,
61 step: 0,
62 m: HashMap::new(),
63 v: HashMap::new(),
64 }
65 }
66
67 /// Override (β₁, β₂).
68 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
69 self.beta1 = b1;
70 self.beta2 = b2;
71 self
72 }
73
74 /// Override the decoupled-decay coefficient.
75 pub fn with_weight_decay(mut self, wd: f32) -> Self {
76 self.weight_decay = wd;
77 self
78 }
79}
80
81impl Optimizer for NAdamW {
82 fn set_lr(&mut self, lr: f32) {
83 self.lr = lr;
84 }
85
86 fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
87 debug_assert_eq!(param.len(), grad.len());
88 let t = (self.step + 1) as f64;
89 let b1 = self.beta1 as f64;
90 let b2 = self.beta2 as f64;
91 let bc1 = 1.0 - b1.powf(t);
92 let bc1_next = 1.0 - b1.powf(t + 1.0);
93 let bc2 = 1.0 - b2.powf(t);
94 let eps = self.eps as f64;
95 let lr = self.lr as f64;
96 let wd = self.weight_decay as f64;
97 let m = zeros_entry(&mut self.m, name, param.len());
98 let v = zeros_entry(&mut self.v, name, param.len());
99 for i in 0..param.len() {
100 let g = grad[i] as f64;
101 let mi = b1 * m[i] as f64 + (1.0 - b1) * g;
102 let vi = b2 * v[i] as f64 + (1.0 - b2) * g * g;
103 m[i] = mi as f32;
104 v[i] = vi as f32;
105 // Nesterov-corrected first moment (Dozat eq. 6):
106 // m_bar = b1 * m_hat_{t+1} + (1-b1)/bc1 * g
107 let m_hat = mi / bc1_next;
108 let m_bar = b1 * m_hat + (1.0 - b1) * g / bc1;
109 let v_hat = vi / bc2;
110 let p = param[i] as f64;
111 param[i] = (p - lr * (m_bar / (v_hat.sqrt() + eps) + wd * p)) as f32;
112 }
113 }
114
115 fn end_iteration(&mut self) {
116 self.step += 1;
117 }
118}