rlx_optim/sophia.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//! Sophia-H — Second-order Clipped Stochastic Optimization (Liu, Xie,
6//! Zhang, Ma, 2023).
7//!
8//! # Idea
9//!
10//! Adam preconditions by `1/√v_t` (a noisy proxy for the inverse
11//! Hessian *diagonal*); Sophia preconditions by the **actual Hessian
12//! diagonal**, computed periodically via a Hutchinson estimator or a
13//! Gauss–Newton approximation. The crucial trick is a *per-coordinate
14//! clip* of the resulting update — even with a noisy Hessian, the
15//! clip caps each coordinate's step at `ρ`, so adversarial curvature
16//! estimates can never blow up the trajectory.
17//!
18//! # Update rule
19//!
20//! ```text
21//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t // first moment EMA
22//! [every K steps the caller updates h via Sophia::update_hessian:]
23//! h ← β₂·h + (1 − β₂)·diag(H_t) // Hessian-diag EMA
24//! u_i = m_{t,i} / max(γ · h_i, ε)
25//! u_i = clip(u_i, −ρ, +ρ)
26//! θ_t = θ_{t-1} − lr · ( u + λ·θ_{t-1} )
27//! ```
28//!
29//! # HVP oracle
30//!
31//! This crate doesn't ship an HVP oracle (it lives in `rlx-autodiff`
32//! as [`rlx_autodiff::hvp`](../../rlx_autodiff/fn.hvp.html)). Call
33//! [`Sophia::update_hessian`] yourself whenever you have a fresh
34//! diagonal estimate (Hutchinson: `H_diag ≈ u ⊙ (∇²L · u)` with random
35//! Rademacher `u`; or Gauss–Newton: `H_diag ≈ g_t²` from a held-out
36//! micro-batch). If you never update it, Sophia degenerates to a
37//! magnitude-clipped first-moment step.
38//!
39//! # When to use
40//!
41//! Curvature-aware optimization for LLM pre-training; the original
42//! paper reports ~2× wall-clock speedup vs AdamW at the same loss.
43//! State cost: two buffers per parameter (`m`, `h`).
44
45use std::collections::HashMap;
46
47use crate::Optimizer;
48use crate::common::zeros_entry;
49
50/// Sophia-H — Hessian-diagonal second-order optimizer.
51#[derive(Debug, Clone)]
52pub struct Sophia {
53 /// Learning rate. Typically slightly *larger* than the AdamW LR
54 /// you'd use on the same model, because the clip bounds the step.
55 pub lr: f32,
56 /// First-moment EMA decay β₁. Default `0.965`.
57 pub beta1: f32,
58 /// Hessian-diagonal EMA decay β₂. Default `0.99`.
59 pub beta2: f32,
60 /// Hessian scale γ (Liu et al. default `0.01`). Multiplies the
61 /// Hessian estimate before forming the denominator.
62 pub gamma: f32,
63 /// Per-coordinate clip threshold ρ. Default `0.04` — the
64 /// dimensionless cap on each step's magnitude.
65 pub rho: f32,
66 /// Denominator floor. Default `1e-12`.
67 pub eps: f32,
68 /// Decoupled weight-decay coefficient λ. Default `0.1` (large by
69 /// AdamW standards — Sophia tolerates more decay).
70 pub weight_decay: f32,
71 step: u64,
72 m: HashMap<String, Vec<f32>>,
73 h: HashMap<String, Vec<f32>>,
74}
75
76impl Sophia {
77 /// Construct with `(β₁, β₂, γ, ρ, ε, λ) = (0.965, 0.99, 0.01, 0.04, 1e-12, 0.1)`.
78 pub fn new(lr: f32) -> Self {
79 Self {
80 lr,
81 beta1: 0.965,
82 beta2: 0.99,
83 gamma: 0.01,
84 rho: 0.04,
85 eps: 1e-12,
86 weight_decay: 0.1,
87 step: 0,
88 m: HashMap::new(),
89 h: HashMap::new(),
90 }
91 }
92
93 /// Override (β₁, β₂).
94 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
95 self.beta1 = b1;
96 self.beta2 = b2;
97 self
98 }
99
100 /// Override the decoupled-decay coefficient.
101 pub fn with_weight_decay(mut self, wd: f32) -> Self {
102 self.weight_decay = wd;
103 self
104 }
105
106 /// Update the diagonal-Hessian estimate for parameter `name`.
107 /// `h_hat` should be a fresh estimate (typically `H_diag` from a
108 /// Hutchinson estimator or `g²` from a Gauss-Newton approximation).
109 pub fn update_hessian(&mut self, name: &str, h_hat: &[f32]) {
110 let h = zeros_entry(&mut self.h, name, h_hat.len());
111 let b2 = self.beta2;
112 for i in 0..h.len() {
113 h[i] = b2 * h[i] + (1.0 - b2) * h_hat[i];
114 }
115 }
116}
117
118impl Optimizer for Sophia {
119 fn set_lr(&mut self, lr: f32) {
120 self.lr = lr;
121 }
122
123 fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
124 debug_assert_eq!(param.len(), grad.len());
125 let b1 = self.beta1;
126 let gamma = self.gamma.max(self.eps);
127 let rho = self.rho;
128 let eps = self.eps;
129 let lr = self.lr;
130 let wd = self.weight_decay;
131 let m = zeros_entry(&mut self.m, name, param.len());
132 for i in 0..param.len() {
133 m[i] = b1 * m[i] + (1.0 - b1) * grad[i];
134 }
135 // Snapshot h (zero if not yet populated).
136 let h_default = vec![0.0f32; param.len()];
137 let h = self.h.get(name).unwrap_or(&h_default);
138 for i in 0..param.len() {
139 let denom = (gamma * h[i]).max(eps);
140 let mut u = m[i] / denom;
141 // Per-coordinate clip to [-rho, rho].
142 if u > rho {
143 u = rho;
144 } else if u < -rho {
145 u = -rho;
146 }
147 // Decoupled decay.
148 param[i] -= lr * (u + wd * param[i]);
149 }
150 }
151
152 fn end_iteration(&mut self) {
153 self.step += 1;
154 }
155}