Skip to main content

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}