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
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)
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