From b3d73a60719db2fe473533b9336751ec438c91cc Mon Sep 17 00:00:00 2001 From: shanyu Date: Fri, 25 Sep 2026 18:13:57 +0800 Subject: [PATCH 1/2] 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/2] 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