From b3d73a60719db2fe473533b9336751ec438c91cc Mon Sep 17 00:00:00 2001 From: shanyu Date: Fri, 25 Sep 2026 18:13:57 +0800 Subject: [PATCH 1/4] Support Conv1d and Conv3d weights in Adahessian trace get_trace averaged the Hessian diagonal over the spatial dims only for 4D kernels. Conv1d (3D) and Conv3d (5D) weights left tmp_output unbound. Average over dims 2..ndim-1 instead; 4D is unchanged. --- tests/test_adahessian_conv.py | 38 +++++++++++++++++++++++++++++++++++ torch_optimizer/adahessian.py | 11 ++++++---- 2 files changed, 45 insertions(+), 4 deletions(-) create mode 100644 tests/test_adahessian_conv.py diff --git a/tests/test_adahessian_conv.py b/tests/test_adahessian_conv.py new file mode 100644 index 00000000..c36c76ad --- /dev/null +++ b/tests/test_adahessian_conv.py @@ -0,0 +1,38 @@ +"""Adahessian must support Conv1d/Conv2d/Conv3d weights (gh issue 500). + +Conv1d weights are 3D and Conv3d weights are 5D. get_trace only averaged +the Hessian diagonal over spatial dims for 4D kernels, so any other conv +rank raised UnboundLocalError. +""" +import torch +import torch_optimizer as optim + + +def _step(module, x): + opt = optim.Adahessian(module.parameters(), lr=0.1) + loss = module(x).pow(2).mean() + loss.backward(create_graph=True) + opt.step() + + +def test_adahessian_conv1d(): + torch.manual_seed(500) + _step(torch.nn.Conv1d(2, 2, 3), torch.randn(1, 2, 8)) + + +def test_adahessian_conv2d(): + torch.manual_seed(500) + _step(torch.nn.Conv2d(2, 2, 3), torch.randn(1, 2, 8, 8)) + + +def test_adahessian_conv3d(): + torch.manual_seed(500) + _step(torch.nn.Conv3d(2, 2, 2), torch.randn(1, 2, 4, 4, 4)) + + +def test_adahessian_conv2d_trace_unchanged(): + # The generic spatial mean must equal the old explicit dims [2, 3]. + hv = torch.randn(2, 2, 3, 3) + expected = torch.mean(hv.abs(), dim=[2, 3], keepdim=True) + got = torch.mean(hv.abs(), dim=list(range(2, len(hv.size()))), keepdim=True) + assert torch.equal(got, expected) diff --git a/torch_optimizer/adahessian.py b/torch_optimizer/adahessian.py index 6c836481..842857a1 100644 --- a/torch_optimizer/adahessian.py +++ b/torch_optimizer/adahessian.py @@ -120,11 +120,14 @@ def get_trace(self, params: Params, grads: Grads) -> List[torch.Tensor]: # We use that torch.abs(hv * vi) = hv.abs() tmp_output = hv.abs() - elif len(param_size) == 4: # Conv kernel - # Hessian diagonal block size is 9 here: torch.sum() reduces - # the dim 2/3. + else: # Conv kernel (Conv1d/Conv2d/Conv3d weights are 3/4/5D) + # Hessian diagonal block sizes are k / k*k / k*k*k here: + # average over the spatial dims, keeping the + # (out_channels, in_channels) block structure. # We use that torch.abs(hv * vi) = hv.abs() - tmp_output = torch.mean(hv.abs(), dim=[2, 3], keepdim=True) + tmp_output = torch.mean( + hv.abs(), dim=list(range(2, len(param_size))), keepdim=True + ) hutchinson_trace.append(tmp_output) return hutchinson_trace From f0b8a43d08dc7bb6bc6d1f859012a9fc342a1dcc Mon Sep 17 00:00:00 2001 From: shanyu Date: Fri, 25 Sep 2026 18:31:33 +0800 Subject: [PATCH 2/4] Raise Shampoo preconditioners to -1/(2k) per the paper Each mode factor of an order-k tensor enters to the power -1/(2k): -1/4 per side for matrices. The code used -1/k, twice the published exponent. The 7 pre-existing test_optimizer Shampoo failures are a state-dict precision issue on this torch version and fail identically without this change. --- tests/test_shampoo_exponent.py | 72 ++++++++++++++++++++++++++++++++++ torch_optimizer/shampoo.py | 7 +++- 2 files changed, 78 insertions(+), 1 deletion(-) create mode 100644 tests/test_shampoo_exponent.py diff --git a/tests/test_shampoo_exponent.py b/tests/test_shampoo_exponent.py new file mode 100644 index 00000000..1fea0a04 --- /dev/null +++ b/tests/test_shampoo_exponent.py @@ -0,0 +1,72 @@ +"""Shampoo must raise each preconditioner to -1/(2k) (gh issue 502). + +Gupta, Koren, Singer 2018 prescribe L^{-1/4} G R^{-1/4} for matrices +(order 2) and L^{-1/(2k)} per mode factor for order-k tensors. The code +used -1/order, i.e. twice the paper exponent. +""" +import torch +import torch_optimizer as optim + + +def _eigh_power(matrix, power): + w, v = torch.linalg.eigh(matrix) + return v @ torch.diag(w.clamp_min(1e-12).pow(power)) @ v.t() + + +def test_matrix_preconditioner_is_inverse_fourth_root(): + # X = M^{-1/4} iff X^4 = M^{-1}. No SVD-vs-eigh sensitivity. + torch.manual_seed(502) + p = torch.nn.Parameter(torch.randn(4, 3)) + opt = optim.Shampoo([p], lr=1e-3, momentum=0, update_freq=1) + p.grad = torch.randn(4, 3) + opt.step() + state = opt.state[p] + for dim_id in range(2): + precond = state["precond_{}".format(dim_id)].cpu() + inv = state["inv_precond_{}".format(dim_id)].cpu() + fourth = inv @ inv @ inv @ inv + assert torch.allclose( + fourth, torch.inverse(precond.double()).float(), + rtol=1e-2, atol=1e-2, + ) + + +def test_vector_preconditioner_is_inverse_sqrt(): + # X = M^{-1/2} iff X^2 = M^{-1}. + torch.manual_seed(502) + p = torch.nn.Parameter(torch.randn(5)) + opt = optim.Shampoo([p], lr=1e-3, momentum=0, update_freq=1) + p.grad = torch.randn(5) + opt.step() + state = opt.state[p] + precond = state["precond_0"].cpu() + inv = state["inv_precond_0"].cpu() + # eps*I keeps one eigenvalue near 1e-4, so the inverse entries reach + # 1e4 and the SVD power carries ~1% relative error. + assert torch.allclose( + inv @ inv, torch.inverse(precond.double()).float(), + rtol=5e-2, atol=1.0, + ) + + +def test_paper_exponent_beats_old_exponent_on_least_squares(): + # Same setup, same lr: the paper exponent (-1/4 per side) must reach a + # lower loss than the old exponent (-1/2 per side). + torch.manual_seed(7) + A = torch.randn(20, 6) + b = torch.randn(20) + + def run(power, steps=60, lr=0.05): + x = torch.zeros(6) + L = torch.eye(6) * 1e-4 + R = torch.eye(1) * 1e-4 + for _ in range(steps): + g = A.T @ (A @ x - b) + G = g.unsqueeze(1) + L = L + G @ G.t() + R = R + G.t() @ G + upd = _eigh_power(L, power) @ G @ _eigh_power(R, power) + x = x - lr * upd.squeeze(1) + return ((A @ x - b) ** 2).mean().item() + + assert run(-0.25) < run(-0.5) diff --git a/torch_optimizer/shampoo.py b/torch_optimizer/shampoo.py index b64c59bc..45c406d9 100644 --- a/torch_optimizer/shampoo.py +++ b/torch_optimizer/shampoo.py @@ -126,7 +126,12 @@ def step(self, closure: OptLossClosure = None) -> OptFloat: grad_t = grad.t() precond.add_(grad @ grad_t) if state["step"] % group["update_freq"] == 0: - inv_precond.copy_(_matrix_power(precond, -1 / order)) + # Gupta, Koren, Singer 2018: each mode-i + # preconditioner enters to the power -1/(2k) for an + # order-k tensor (-1/4 per side for matrices). + inv_precond.copy_( + _matrix_power(precond, -1 / (2 * order)) + ) if dim_id == order - 1: # finally From 1cf34c32f3dd1f03d4f763498035ebc92fca262c Mon Sep 17 00:00:00 2001 From: shanyu Date: Fri, 25 Sep 2026 18:55:05 +0800 Subject: [PATCH 3/4] Make Adafactor relative-step constants configurable warmup_rate (default 1e-6) replaces the hard-coded warm-up slope and min_step_size (default 1e-2) the floor. Old checkpoints without the new group keys fall back to the old values. The 7 pre-existing test_optimizer Adafactor failures are a state-dict precision issue on this torch version and fail identically without this change. --- tests/test_adafactor_step_size.py | 60 +++++++++++++++++++++++++++++++ torch_optimizer/adafactor.py | 24 +++++++++++-- 2 files changed, 82 insertions(+), 2 deletions(-) create mode 100644 tests/test_adafactor_step_size.py diff --git a/tests/test_adafactor_step_size.py b/tests/test_adafactor_step_size.py new file mode 100644 index 00000000..e26bf6f4 --- /dev/null +++ b/tests/test_adafactor_step_size.py @@ -0,0 +1,60 @@ +"""Adafactor relative-step constants are configurable (gh issue 535).""" +import math + +import pytest +import torch +import torch_optimizer as optim + + +def _lr(opt, step, rms=1.0): + group = opt.param_groups[0] + state = {"step": step, "RMS": rms} + return opt._get_lr(group, state) + + +def test_defaults_reproduce_old_schedule(): + p = torch.nn.Parameter(torch.ones(4)) + opt = optim.Adafactor([p], warmup_init=True) + assert _lr(opt, 100) == min(1e-6 * 100, 1.0 / math.sqrt(100)) + opt2 = optim.Adafactor([p], warmup_init=False) + assert _lr(opt2, 100) == min(1e-2, 1.0 / math.sqrt(100)) + + +def test_custom_warmup_rate_scales_early_steps(): + p = torch.nn.Parameter(torch.ones(4)) + opt = optim.Adafactor([p], warmup_init=True, warmup_rate=2e-6) + assert _lr(opt, 100) == min(2e-6 * 100, 1.0 / math.sqrt(100)) + assert _lr(opt, 100) == 2 * min(1e-6 * 100, 1.0 / math.sqrt(100)) + + +def test_custom_min_step_size(): + p = torch.nn.Parameter(torch.ones(4)) + opt = optim.Adafactor([p], warmup_init=False, min_step_size=5e-3) + assert _lr(opt, 100) == min(5e-3, 1.0 / math.sqrt(100)) + + +def test_old_checkpoint_without_new_keys(): + p = torch.nn.Parameter(torch.ones(4)) + opt = optim.Adafactor([p]) + del opt.param_groups[0]["warmup_rate"] + del opt.param_groups[0]["min_step_size"] + assert _lr(opt, 100) == min(1e-2, 1.0 / math.sqrt(100)) + + +def test_negative_values_rejected(): + p = torch.nn.Parameter(torch.ones(4)) + with pytest.raises(ValueError): + optim.Adafactor([p], warmup_rate=-1e-6) + with pytest.raises(ValueError): + optim.Adafactor([p], min_step_size=-1e-2) + + +def test_end_to_end_step(): + torch.manual_seed(535) + model = torch.nn.Linear(4, 2) + opt = optim.Adafactor( + model.parameters(), warmup_init=True, warmup_rate=2e-6 + ) + loss = model(torch.randn(3, 4)).pow(2).mean() + loss.backward() + opt.step() diff --git a/torch_optimizer/adafactor.py b/torch_optimizer/adafactor.py index ed1756b9..e26b5a48 100644 --- a/torch_optimizer/adafactor.py +++ b/torch_optimizer/adafactor.py @@ -35,6 +35,11 @@ class Adafactor(Optimizer): instead of external learning rate (default: True) warmup_init: time-dependent learning rate computation depends on whether warm-up initialization is being used (default: False) + warmup_rate: slope of the warm-up ramp, used as + ``warmup_rate * step`` when ``warmup_init`` is true + (default: 1e-6) + min_step_size: floor of the relative step size when ``warmup_init`` + is false (default: 1e-2) Example: >>> import torch_optimizer as optim @@ -61,6 +66,8 @@ def __init__( scale_parameter: bool = True, relative_step: bool = True, warmup_init: bool = False, + warmup_rate: float = 1e-6, + min_step_size: float = 1e-2, ): if lr is not None and lr <= 0.0: raise ValueError("Invalid learning rate: {}".format(lr)) @@ -68,6 +75,14 @@ def __init__( raise ValueError( "Invalid weight_decay value: {}".format(weight_decay) ) + if warmup_rate < 0.0: + raise ValueError( + "Invalid warmup_rate value: {}".format(warmup_rate) + ) + if min_step_size < 0.0: + raise ValueError( + "Invalid min_step_size value: {}".format(min_step_size) + ) defaults = dict( lr=lr, @@ -79,16 +94,21 @@ def __init__( scale_parameter=scale_parameter, relative_step=relative_step, warmup_init=warmup_init, + warmup_rate=warmup_rate, + min_step_size=min_step_size, ) super(Adafactor, self).__init__(params, defaults) def _get_lr(self, param_group: ParamGroup, param_state: State) -> float: + # Groups saved before warmup_rate/min_step_size existed lack them. + param_group.setdefault("warmup_rate", 1e-6) + param_group.setdefault("min_step_size", 1e-2) rel_step_sz = param_group["lr"] if param_group["relative_step"]: min_step = ( - 1e-6 * param_state["step"] + param_group["warmup_rate"] * param_state["step"] if param_group["warmup_init"] - else 1e-2 + else param_group["min_step_size"] ) rel_step_sz = min(min_step, 1.0 / math.sqrt(param_state["step"])) param_scale = 1.0 From 0572c5c414bff8474541fb64de2b33d01b2c892d Mon Sep 17 00:00:00 2001 From: shanyu Date: Fri, 25 Sep 2026 19:06:30 +0800 Subject: [PATCH 4/4] Add the Sophia optimizer Second-order optimizer with a diagonal Hessian estimate (SophiaG). Follows the reference update, including update_hessian refresh, maximize flag, and sparse-gradient rejection. CUDA-graph capturable execution is not ported. Wired into the test harness lists and README. --- README.rst | 30 ++++++ tests/test_optimizer.py | 1 + tests/test_optimizer_with_nn.py | 1 + tests/test_param_validation.py | 2 + tests/test_sophia.py | 77 ++++++++++++++ torch_optimizer/__init__.py | 2 + torch_optimizer/sophia.py | 183 ++++++++++++++++++++++++++++++++ 7 files changed, 296 insertions(+) create mode 100644 tests/test_sophia.py create mode 100644 torch_optimizer/sophia.py diff --git a/README.rst b/README.rst index bc7d5b17..bc598792 100644 --- a/README.rst +++ b/README.rst @@ -149,6 +149,9 @@ Supported Optimizers | `Shampoo`_ | https://arxiv.org/abs/1802.09568 | +---------------+--------------------------------------------------------------------------------------------------------------------------------------+ | | | +| `Sophia`_ | https://arxiv.org/abs/2305.14342 | ++---------------+--------------------------------------------------------------------------------------------------------------------------------------+ +| | | | `Yogi`_ | https://papers.nips.cc/paper/8186-adaptive-methods-for-nonconvex-optimization | +---------------+--------------------------------------------------------------------------------------------------------------------------------------+ @@ -1007,6 +1010,33 @@ Shampoo **Reference Code**: https://github.com/moskomule/shampoo.pytorch +Sophia +------ + +.. code:: python + + import torch_optimizer as optim + + # model = ... + optimizer = optim.Sophia( + m.parameters(), + lr=1e-4, + betas=(0.965, 0.99), + rho=0.04, + weight_decay=1e-1, + ) + optimizer.zero_grad() + loss_fn(model(input), target).backward() + optimizer.update_hessian() + optimizer.zero_grad() + loss_fn(model(input), target).backward() + optimizer.step() + +**Paper**: *Sophia: A Scalable Stochastic Second-order Optimizer for Language Model Pre-training* (2023) [https://arxiv.org/abs/2305.14342] + +**Reference Code**: https://github.com/Liuhong99/Sophia + + Yogi ---- diff --git a/tests/test_optimizer.py b/tests/test_optimizer.py index 154474e2..369da635 100644 --- a/tests/test_optimizer.py +++ b/tests/test_optimizer.py @@ -110,6 +110,7 @@ def build_lookahead(*a, **kw): optim.Shampoo, optim.Yogi, optim.Lion, + optim.Sophia, ] diff --git a/tests/test_optimizer_with_nn.py b/tests/test_optimizer_with_nn.py index 80829ad0..00e29d25 100644 --- a/tests/test_optimizer_with_nn.py +++ b/tests/test_optimizer_with_nn.py @@ -90,6 +90,7 @@ def build_lookahead(*a, **kw): (optim.Yogi, {"lr": 0.1, "weight_decay": 1e-3}, 200), (optim.Adahessian, {"lr": 0.1, "weight_decay": 1e-3}, 200), (optim.Lion, {"lr": 0.1, "weight_decay": 1e-3}, 200), + (optim.Sophia, {"lr": 0.1, "weight_decay": 1e-3}, 200), ] diff --git a/tests/test_param_validation.py b/tests/test_param_validation.py index c5d74390..e091716f 100644 --- a/tests/test_param_validation.py +++ b/tests/test_param_validation.py @@ -56,6 +56,7 @@ def test_sparse_not_supported(optimizer_class): optim.Shampoo, optim.Yogi, optim.Lion, + optim.Sophia, ] @@ -120,6 +121,7 @@ def test_eps_validation(optimizer_class): optim.Shampoo, optim.Yogi, optim.Lion, + optim.Sophia, ] diff --git a/tests/test_sophia.py b/tests/test_sophia.py new file mode 100644 index 00000000..6b5c5c3d --- /dev/null +++ b/tests/test_sophia.py @@ -0,0 +1,77 @@ +"""Sophia optimizer: matches the reference update, Hessian refresh works.""" +import torch +import torch_optimizer as optim + + +def test_matches_reference_update(): + torch.manual_seed(1) + p = torch.nn.Parameter(torch.randn(3, 4)) + opt = optim.Sophia([p], lr=1e-3, betas=(0.9, 0.99), rho=0.04, + weight_decay=0.0) + g = torch.randn(3, 4) + h = torch.rand(3, 4) + 0.5 + p.grad = g.clone() + # seed state as update_hessian would after one refresh + opt.update_hessian() + opt.state[p]["hessian"].copy_(h) + opt.state[p]["exp_avg"].zero_() + before = p.detach().clone() + opt.step() + beta1, bs, rho, lr = 0.9, 5120, 0.04, 1e-3 + m = (1 - beta1) * g + ratio = (m.abs() / (rho * bs * h + 1e-15)).clamp(None, 1) + expected = before * (1 - lr * 0.0) - lr * m.sign() * ratio + assert torch.allclose(p.detach(), expected, atol=1e-6) + + +def test_update_hessian_accumulates(): + torch.manual_seed(2) + p = torch.nn.Parameter(torch.randn(4)) + opt = optim.Sophia([p]) + p.grad = torch.ones(4) * 2.0 + opt.update_hessian() + h1 = opt.state[p]["hessian"].clone() + assert torch.allclose(h1, torch.full((4,), (1 - 0.99) * 4.0)) + opt.update_hessian() + h2 = opt.state[p]["hessian"].clone() + assert torch.allclose(h2, h1 * 0.99 + (1 - 0.99) * 4.0) + + +def test_step_without_hessian_update_is_signed_step(): + torch.manual_seed(3) + p = torch.nn.Parameter(torch.zeros(4)) + opt = optim.Sophia([p], lr=0.01, weight_decay=0.0) + p.grad = torch.tensor([1.0, -2.0, 0.5, -0.25]) + opt.step() + # hessian is zero -> ratio clamps to 1 -> pure sign step + assert torch.allclose( + p.detach(), -0.01 * torch.tensor([1.0, -1.0, 1.0, -1.0]) + ) + + +def test_sparse_rejected(): + p = torch.nn.Parameter(torch.randn(4)) + opt = optim.Sophia([p]) + p.grad = torch.sparse_coo_tensor( + torch.tensor([[0, 2]]), torch.tensor([1.0, 2.0]), (4,) + ) + try: + opt.step() + except RuntimeError: + return + raise AssertionError("sparse gradient was not rejected") + + +def test_maximize_flips_sign(): + torch.manual_seed(4) + kw = dict(lr=0.01, weight_decay=0.0) + a = torch.nn.Parameter(torch.zeros(4)) + b = torch.nn.Parameter(torch.zeros(4)) + oa = optim.Sophia([a], **kw) + ob = optim.Sophia([b], maximize=True, **kw) + g = torch.tensor([1.0, 2.0, 3.0, 4.0]) + a.grad = g.clone() + b.grad = g.clone() + oa.step() + ob.step() + assert torch.allclose(a.detach(), -b.detach()) diff --git a/torch_optimizer/__init__.py b/torch_optimizer/__init__.py index f0123e53..b54510d7 100644 --- a/torch_optimizer/__init__.py +++ b/torch_optimizer/__init__.py @@ -43,6 +43,7 @@ from .sgdp import SGDP from .sgdw import SGDW from .shampoo import Shampoo +from .sophia import Sophia from .swats import SWATS from .yogi import Yogi @@ -76,6 +77,7 @@ "SGDW", "SWATS", "Shampoo", + "Sophia", "Yogi", "Lion", # utils diff --git a/torch_optimizer/sophia.py b/torch_optimizer/sophia.py new file mode 100644 index 00000000..1abc4b81 --- /dev/null +++ b/torch_optimizer/sophia.py @@ -0,0 +1,183 @@ +import torch +from torch.optim.optimizer import Optimizer + +from .types import Betas2, OptFloat, OptLossClosure, Params + +__all__ = ("Sophia",) + + +class Sophia(Optimizer): + r"""Implements SophiaG algorithm. + + Sophia - A Scalable Stochastic Second-order Optimizer for Language Model + Pre-training. It uses a diagonal Hessian estimate to scale the update of + a momentum-normalized gradient, which allows larger step sizes than + first-order methods on language models. + + The Hessian estimate must be refreshed with :meth:`update_hessian` + (usually on a sampled batch every few steps) before calling + :meth:`step`. Without any Hessian update the estimate stays zero and + the update degrades to a signed momentum step. + + This port follows the non-capturable path of the reference + implementation. CUDA-graph capturable execution is not supported. + + Arguments: + params: iterable of parameters to optimize or dicts defining + parameter groups + lr: learning rate (default: 1e-4) + betas: coefficients used for computing running averages of gradient + and Hessian diagonal (default: (0.965, 0.99)) + rho: Hutchinson trace clipping factor (default: 0.04) + weight_decay: weight decay (L2 penalty) (default: 1e-1) + maximize: maximize the params based on the objective, instead of + minimizing (default: False) + + Example: + >>> import torch_optimizer as optim + >>> optimizer = optim.Sophia(model.parameters(), lr=1e-4) + >>> optimizer.zero_grad() + >>> loss_fn(model(input), target).backward() + >>> optimizer.update_hessian() + >>> optimizer.zero_grad() + >>> loss_fn(model(input), target).backward() + >>> optimizer.step() + + __ https://arxiv.org/abs/2305.14342 + + Note: + Reference code: https://github.com/Liuhong99/Sophia + """ + + def __init__( + self, + params: Params, + lr: float = 1e-4, + betas: Betas2 = (0.965, 0.99), + rho: float = 0.04, + weight_decay: float = 1e-1, + maximize: bool = False, + ) -> None: + if not 0.0 <= lr: + raise ValueError("Invalid learning rate: {}".format(lr)) + if not 0.0 <= betas[0] < 1.0: + raise ValueError( + "Invalid beta parameter at index 0: {}".format(betas[0]) + ) + if not 0.0 <= betas[1] < 1.0: + raise ValueError( + "Invalid beta parameter at index 1: {}".format(betas[1]) + ) + if not 0.0 <= rho: + raise ValueError("Invalid rho value: {}".format(rho)) + if not 0.0 <= weight_decay: + raise ValueError( + "Invalid weight_decay value: {}".format(weight_decay) + ) + defaults = dict( + lr=lr, + betas=betas, + rho=rho, + weight_decay=weight_decay, + maximize=maximize, + ) + super(Sophia, self).__init__(params, defaults) + + def __setstate__(self, state) -> None: + super(Sophia, self).__setstate__(state) + for group in self.param_groups: + group.setdefault("maximize", False) + + @torch.no_grad() + def update_hessian(self) -> None: + """Refresh the diagonal Hessian estimate from current gradients.""" + for group in self.param_groups: + _, beta2 = group["betas"] + for p in group["params"]: + if p.grad is None: + continue + state = self.state[p] + if len(state) == 0: + state["step"] = torch.tensor(0.0) + state["exp_avg"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + state["hessian"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + if "hessian" not in state.keys(): + state["hessian"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + state["hessian"].mul_(beta2).addcmul_( + p.grad, p.grad, value=1 - beta2 + ) + + @torch.no_grad() + def step( + self, closure: OptLossClosure = None, bs: int = 5120 + ) -> OptFloat: + """Perform a single optimization step. + + Arguments: + closure: A closure that reevaluates the model and returns + the loss. + bs: batch size used to scale the Hessian estimate + (default: 5120). + """ + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + + for group in self.param_groups: + beta1, _ = group["betas"] + for p in group["params"]: + if p.grad is None: + continue + grad = p.grad + if grad.is_sparse: + raise RuntimeError( + "Sophia does not support sparse gradients" + ) + state = self.state[p] + if len(state) == 0: + state["step"] = torch.tensor(0.0) + state["exp_avg"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + state["hessian"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + if "hessian" not in state.keys(): + state["hessian"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + exp_avg = state["exp_avg"] + hess = state["hessian"] + + if torch.is_complex(p): + grad = torch.view_as_real(grad) + exp_avg = torch.view_as_real(exp_avg) + hess = torch.view_as_real(hess) + p_view = torch.view_as_real(p) + else: + p_view = p + + if group["maximize"]: + grad = -grad + + state["step"] += 1 + + # Perform stepweight decay + p_view.mul_(1 - group["lr"] * group["weight_decay"]) + + # Decay the first moment running average coefficient + exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1) + + ratio = ( + exp_avg.abs() / (group["rho"] * bs * hess + 1e-15) + ).clamp(None, 1) + p_view.addcmul_(exp_avg.sign(), ratio, value=-group["lr"]) + + return loss