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
57 changes: 57 additions & 0 deletions README.rst
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,9 @@ Supported Optimizers
| `AccSGD`_ | https://arxiv.org/abs/1803.05591 |
+---------------+--------------------------------------------------------------------------------------------------------------------------------------+
| | |
| `Aida`_ | https://arxiv.org/abs/2203.13273 |
+---------------+--------------------------------------------------------------------------------------------------------------------------------------+
| | |
| `AdaBelief`_ | https://arxiv.org/abs/2010.07468 |
+---------------+--------------------------------------------------------------------------------------------------------------------------------------+
| | |
Expand Down Expand Up @@ -149,6 +152,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 @@ -304,6 +310,30 @@ AccSGD
**Reference Code**: https://github.com/rahulkidambi/AccSGD


Aida
----

.. code:: python

import torch_optimizer as optim

# model = ...
optimizer = optim.Aida(
m.parameters(),
lr=1e-3,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=0,
k=2,
xi=1e-20,
)
optimizer.step()

**Paper**: *On Exploiting Layerwise Gradient Statistics for Effective Training of Deep Neural Networks* (2022) [https://arxiv.org/abs/2203.13273]

**Reference Code**: https://github.com/guoqiang-zhang-x/Aida-Optimizer


AdaBelief
---------

Expand Down Expand Up @@ -1007,6 +1037,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)
88 changes: 88 additions & 0 deletions tests/test_aida.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
"""Aida optimizer: matches the reference mathematics, no eps drift."""
import math

import torch
import torch_optimizer as optim


def _reference_step(p, g, state, lr, beta1, beta2, eps, wd, k, xi):
m, v, step = state
step += 1
bc1 = 1 - beta1**step
bc2 = 1 - beta2**step
g = g + wd * p
m = beta1 * m + (1 - beta1) * g
pg, pm = g.clone(), m.clone()
for _ in range(k):
inner = (pg * pm).sum()
pg = pg * (inner / ((pg**2).sum() + xi))
pm = pm * (inner / ((pm**2).sum() + xi))
r = pm - pg
v = beta2 * v + (1 - beta2) * r**2
denom = (v + eps).sqrt() / math.sqrt(bc2)
p = p - (lr / bc1) * m / denom
return p, (m, v, step)


def test_matches_reference_two_steps():
torch.manual_seed(9)
for k in (1, 2, 3):
p = torch.nn.Parameter(torch.randn(6))
opt = optim.Aida([p], lr=1e-3, weight_decay=1e-2, k=k)
state = (
torch.zeros(6),
torch.zeros(6),
0,
)
ref_p = p.detach().clone()
for _ in range(2):
g = torch.randn(6)
p.grad = g.clone()
opt.step()
ref_p, state = _reference_step(
ref_p, g, state, 1e-3, 0.9, 0.999, 1e-8, 1e-2, k, 1e-20
)
assert torch.allclose(p.detach(), ref_p, atol=1e-6)


def test_eps_does_not_accumulate_in_state():
# After one step from zero state, exp_avg_var must be exactly
# (1-b2) * residual^2 with no eps mixed in. The reference added eps
# into the state in place, so eps would already pollute it here.
torch.manual_seed(11)
g = torch.randn(4)
p = torch.nn.Parameter(torch.zeros(4))
opt = optim.Aida([p], lr=1e-3, eps=0.5)
p.grad = g.clone()
opt.step()
m = (1 - 0.9) * g
pg, pm = g.clone(), m.clone()
for _ in range(2):
inner = (pg * pm).sum()
pg = pg * (inner / ((pg**2).sum() + 1e-20))
pm = pm * (inner / ((pm**2).sum() + 1e-20))
expected = (1 - 0.999) * (pm - pg) ** 2
assert torch.allclose(opt.state[p]["exp_avg_var"], expected, atol=1e-9)


def test_reset_clears_state():
p = torch.nn.Parameter(torch.randn(4))
opt = optim.Aida([p])
p.grad = torch.randn(4)
opt.step()
assert opt.state[p]["step"] == 1
opt.reset()
assert opt.state[p]["step"] == 0
assert torch.equal(
opt.state[p]["exp_avg"], torch.zeros_like(p)
)


def test_invalid_k_xi_rejected():
p = torch.nn.Parameter(torch.ones(2))
import pytest

with pytest.raises(ValueError):
optim.Aida([p], k=0)
with pytest.raises(ValueError):
optim.Aida([p], xi=0.0)
2 changes: 2 additions & 0 deletions tests/test_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,7 @@ def build_lookahead(*a, **kw):


optimizers = [
optim.Aida,
build_lookahead,
optim.A2GradExp,
optim.A2GradInc,
Expand Down Expand Up @@ -110,6 +111,7 @@ def build_lookahead(*a, **kw):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
]


Expand Down
2 changes: 2 additions & 0 deletions tests/test_optimizer_with_nn.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ def build_lookahead(*a, **kw):

optimizers = [
(build_lookahead, {"lr": 0.1, "weight_decay": 1e-3}, 200),
(optim.Aida, {"lr": 0.1, "weight_decay": 1e-3}, 200),
(optim.A2GradExp, {"lips": 2.0, "beta": 1e-3}, 500),
(optim.A2GradInc, {"lips": 5.0, "beta": 1e-3}, 200),
(optim.A2GradUni, {"lips": 5.0, "beta": 1e-3}, 500),
Expand Down Expand Up @@ -90,6 +91,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
4 changes: 4 additions & 0 deletions tests/test_param_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,8 @@ def test_sparse_not_supported(optimizer_class):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
optim.Aida,
]


Expand Down Expand Up @@ -98,6 +100,7 @@ def test_eps_validation(optimizer_class):

weight_decay_optimizers = [
optim.AccSGD,
optim.Aida,
optim.AdaBelief,
optim.AdaBound,
optim.AdaMod,
Expand All @@ -120,6 +123,7 @@ def test_eps_validation(optimizer_class):
optim.Shampoo,
optim.Yogi,
optim.Lion,
optim.Sophia,
]


Expand Down
Loading