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
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)
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)
24 changes: 22 additions & 2 deletions torch_optimizer/adafactor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -61,13 +66,23 @@ 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))
if weight_decay < 0.0:
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,
Expand All @@ -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
Expand Down
11 changes: 7 additions & 4 deletions torch_optimizer/adahessian.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 6 additions & 1 deletion torch_optimizer/shampoo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down