Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions README.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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 |
+---------------+--------------------------------------------------------------------------------------------------------------------------------------+

Expand Down Expand Up @@ -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
----

Expand Down
60 changes: 60 additions & 0 deletions tests/test_adafactor_step_size.py
Original file line number Diff line number Diff line change
@@ -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()
38 changes: 38 additions & 0 deletions tests/test_adahessian_conv.py
Original file line number Diff line number Diff line change
@@ -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)
1 change: 1 addition & 0 deletions tests/test_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,7 @@ def build_lookahead(*a, **kw):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
]


Expand Down
1 change: 1 addition & 0 deletions tests/test_optimizer_with_nn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
]


Expand Down
2 changes: 2 additions & 0 deletions tests/test_param_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ def test_sparse_not_supported(optimizer_class):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
]


Expand Down Expand Up @@ -120,6 +121,7 @@ def test_eps_validation(optimizer_class):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
]


Expand Down
72 changes: 72 additions & 0 deletions tests/test_shampoo_exponent.py
Original file line number Diff line number Diff line change
@@ -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)
77 changes: 77 additions & 0 deletions tests/test_sophia.py
Original file line number Diff line number Diff line change
@@ -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())
2 changes: 2 additions & 0 deletions torch_optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -76,6 +77,7 @@
"SGDW",
"SWATS",
"Shampoo",
"Sophia",
"Yogi",
"Lion",
# utils
Expand Down
Loading