From 41a4879820a3be7ff3c0437e07b081b5ac147ceb Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 10 Aug 2026 13:55:01 -0700 Subject: [PATCH 01/16] Use type error for shape mismatch Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/soap_v3.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/emerging_optimizers/shampoo/soap_v3.py b/emerging_optimizers/shampoo/soap_v3.py index 58a6695..c52b1c3 100644 --- a/emerging_optimizers/shampoo/soap_v3.py +++ b/emerging_optimizers/shampoo/soap_v3.py @@ -78,7 +78,7 @@ def init_state( ValueError: If ``shape`` is not 2D. """ if len(shape) != 2: - raise ValueError(f"KlSoapPreconditioner is only supported for 2D tensors, got shape {tuple(shape)}") + raise TypeError(f"KlSoapPreconditioner is only supported for 2D tensors, got shape {tuple(shape)}") m, n = shape return { "exp_avg": torch.zeros(m, n, device=device), From fe48cae644263a8c1b7c0140675bc25710f71b5d Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 10 Aug 2026 19:36:15 -0700 Subject: [PATCH 02/16] add vanilla shampoo Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 308 ++++++++++++++++++++ emerging_optimizers/shampoo/shampoo_base.py | 4 + 2 files changed, 312 insertions(+) create mode 100644 emerging_optimizers/shampoo/shampoo.py diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py new file mode 100644 index 0000000..2d17701 --- /dev/null +++ b/emerging_optimizers/shampoo/shampoo.py @@ -0,0 +1,308 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import TYPE_CHECKING, Callable, ClassVar, override + + +if TYPE_CHECKING: + from typing import overload + +import torch +from torch import optim +from torch.optim.optimizer import ParamsT + +from emerging_optimizers import mixin as opt_mixin +from emerging_optimizers.shampoo import shampoo_base +from emerging_optimizers.utils import eig as eig_utils + + +__all__ = [ + "Shampoo", + "ShampooPreconditioner", +] + + +class ShampooPreconditioner: + """Per-parameter Shampoo preconditioner holding the Kronecker factors of one 2D parameter. + + Args: + state: Per-parameter optimizer state holding ``L`` and ``R``. + p_inv_root: Inverse root order; each factor is applied as ``A^(-1/p_inv_root)``. + eps: Floor on the eigenvalue magnitudes before inversion. + """ + + def __init__( + self, + state: dict, + p_inv_root: float, + eps: float = 1e-8, + ) -> None: + self.kronecker_factor_pair = shampoo_base.TensorPair(state["L"], state["R"]) + self.p_inv_root = p_inv_root + self.eps = eps + + @staticmethod + def init_state( + shape: tuple[int, ...], + device: torch.device, + ) -> dict[str, torch.Tensor]: + """Creates the Kronecker factors for a parameter shape. + + Args: + shape: Shape of the 2D parameter the preconditioner will be attached to. + device: Device to allocate the state tensors on. + + Returns: + The state entries owned by this preconditioner, keyed as :meth:`rebind_state` expects them. + + Raises: + TypeError: If ``shape`` is not 2D. + """ + if len(shape) != 2: + raise TypeError(f"ShampooPreconditioner is only supported for 2D tensors, got shape {tuple(shape)}") + m, n = shape + return { + "L": torch.zeros(m, m, device=device), + "R": torch.zeros(n, n, device=device), + } + + def rebind_state(self, state: dict) -> None: + """Binds the current preconditioner tensors back into the optimizer state dict. + + Args: + state: Per-parameter optimizer state, updated in place. + + Raises: + KeyError: If ``state`` is missing any of the preconditioner keys. + """ + updates = { + "L": self.kronecker_factor_pair.L, + "R": self.kronecker_factor_pair.R, + } + missing = updates.keys() - state.keys() + if missing: + raise KeyError(f"rebind_state: state missing keys {sorted(missing)}") + state.update(updates) + + def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: + """Performs the first step's factor update, before any history exists. + + Args: + grad: Gradient of the parameter on the first step. + shampoo_beta: EMA coefficient for the Kronecker factor update. + """ + self.update_kronecker_factors(grad, shampoo_beta) + + def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: + """Accumulates the gradient outer products into the Kronecker factors. + + Args: + grad: Gradient of the parameter. + shampoo_beta: EMA coefficient for the Kronecker factor update. + """ + self.kronecker_factor_pair.L.lerp_(grad @ grad.T, 1 - shampoo_beta) + self.kronecker_factor_pair.R.lerp_(grad.T @ grad, 1 - shampoo_beta) + + def step(self, grad: torch.Tensor, shampoo_beta: float) -> None: + """Updates the Kronecker factors with the latest gradient. + + Args: + grad: Gradient of the parameter. + shampoo_beta: EMA coefficient for the Kronecker factor update. + """ + self.update_kronecker_factors(grad, shampoo_beta) + + def _get_inverse_root(self, kronecker_factor: torch.Tensor) -> torch.Tensor: + """Computes ``kronecker_factor^(-1/p_inv_root)`` from its eigendecomposition. + + Args: + kronecker_factor: left or right kronecker factor + + Returns: + The inverse root of the factor. + """ + eigvals, eigvecs = eig_utils.eigh_with_fallback(kronecker_factor) + # Kronecker factors are symmetric. Eigh can sometime return negative values for numerical 0 therefore the abs. + return (eigvecs * eigvals.abs().clamp_min(self.eps) ** (-1.0 / self.p_inv_root)) @ eigvecs.mT + + def precondition(self, x: torch.Tensor) -> torch.Tensor: + """Applies both inverse roots to a matrix in the parameter basis. + + Args: + x: Matrix in the parameter basis. + + Returns: + The preconditioned matrix, in the parameter basis. + """ + inverse_root_pair = shampoo_base.TensorPair( + self._get_inverse_root(self.kronecker_factor_pair.L), + self._get_inverse_root(self.kronecker_factor_pair.R), + ) + + return inverse_root_pair.L @ x @ inverse_root_pair.R + + +class Shampoo(optim.Optimizer, opt_mixin.WeightDecayMixin): + """Shampoo with inverse roots rebuilt by eigendecomposition on every step. + + The update is EMA momentum in the parameter basis, preconditioned on both sides by the inverse roots + of the Kronecker factors. + + Args: + params: Iterable of 2D CUDA parameters to optimize or dicts defining parameter groups. + lr: Learning rate. + momentum: Momentum EMA coefficient. + shampoo_beta: Kronecker factor EMA coefficient. + weight_decay: Decoupled weight decay coefficient. + p_inv_root: Inverse root order applied to each Kronecker factor. + + Attributes: + PreconditionerCls: Preconditioner used for every parameter, and the source of the state layout + allocated by :meth:`_init_group`. Subclasses set it to change how the factors are maintained. + """ + + PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] = ShampooPreconditioner + + def __init__( + self, + params: ParamsT, + lr: float, + momentum: float = 0.9, + shampoo_beta: float = 0.95, + weight_decay: float = 0.01, + *, + p_inv_root: float = 4, + ) -> None: + self.weight_decay_method = "decoupled" + self.p_inv_root = p_inv_root + + if lr < 0.0: + raise ValueError(f"Invalid learning rate: {lr}") + + defaults = { + "lr": lr, + "momentum": momentum, + "shampoo_beta": shampoo_beta, + "weight_decay": weight_decay, + } + super().__init__(params, defaults) + + @torch.compile + def _scalar_update( + self, + grad: torch.Tensor, + exp_avg: torch.Tensor, + *, + momentum: float, + ) -> torch.Tensor: + """Applies the inner scalar optimizer to the gradient, in the parameter basis. + + Args: + grad: Gradient of the parameter. + exp_avg: Momentum buffer, updated in place. + momentum: Momentum EMA coefficient. + + Returns: + The scalar update, in the parameter basis. + """ + exp_avg.lerp_(grad, 1 - momentum) + return exp_avg + + @torch.no_grad() # type: ignore[misc] + def _init_group( + self, + group: dict, + skip_non_grad_params: bool = True, + ) -> None: + """Performs lazy state initialization for parameters with gradients. + + Args: + group: Parameter group dictionary. + skip_non_grad_params: Whether to skip parameters with no gradients. + + Raises: + TypeError: If the parameter is not a 2D CUDA tensor. + """ + for p in group["params"]: + if skip_non_grad_params and p.grad is None: + continue + + if p.dim() != 2: + raise TypeError(f"{type(self).__name__} is only supported for 2D tensors") + if not p.is_cuda: + raise TypeError(f"{type(self).__name__} only supports CUDA tensors") + + state = self.state[p] + + if len(state) == 0: + state["step"] = 0 + state["exp_avg"] = torch.zeros_like(p, dtype=torch.float32) + + state.update(self.PreconditionerCls.init_state(p.shape, p.device)) + + if TYPE_CHECKING: + + @overload + def step(self, closure: None = ...) -> None: ... + + @overload + def step(self, closure: Callable[[], float]) -> float: ... + + @torch.no_grad() # type: ignore[misc] + @override + def step(self, closure: Callable[[], float] | None = None) -> float | None: + """Performs a single optimization step. + + Args: + closure: Unsupported; must be ``None``. + + Raises: + ValueError: If ``closure`` is not ``None``. + """ + if closure is not None: + raise ValueError("closure is not supported") + + for group in self.param_groups: + self._init_group(group) + + for group in self.param_groups: + for p in group["params"]: + if p.grad is None: + continue # pragma: no cover + + grad = p.grad.to(torch.float32) + state = self.state[p] + + preconditioner = self.PreconditionerCls(state, self.p_inv_root) + + scalar_update = self._scalar_update(grad, state["exp_avg"], momentum=group["momentum"]) + + if state["step"] == 0: + preconditioner.init_step(grad, group["shampoo_beta"]) + else: + preconditioner.step(grad, group["shampoo_beta"]) + preconditioned_update = preconditioner.precondition(scalar_update) + + self._apply_weight_decay_inplace( + p, + grad, + group["lr"], + group["weight_decay"], + ) + p.add_(preconditioned_update.to(p.dtype), alpha=-group["lr"]) + + preconditioner.rebind_state(state) + state["step"] += 1 + + return None diff --git a/emerging_optimizers/shampoo/shampoo_base.py b/emerging_optimizers/shampoo/shampoo_base.py index 77dbcb1..5f9e675 100644 --- a/emerging_optimizers/shampoo/shampoo_base.py +++ b/emerging_optimizers/shampoo/shampoo_base.py @@ -37,6 +37,10 @@ def __iter__(self) -> Iterator[torch.Tensor]: """Iterates over the pair as ``L`` then ``R``.""" return iter((self.L, self.R)) + def __getitem__(self, index: int) -> torch.Tensor: + """Indexes the pair as ``0`` for ``L`` and ``1`` for ``R``.""" + return (self.L, self.R)[index] + class _PreconditionerProtocol(Protocol): """Interface every preconditioner in the family must provide, for one parameter. From 4ae48bdd1e1b44827c5dab1b2d6350cbfdd56c19 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 13 Aug 2026 09:34:58 -0700 Subject: [PATCH 03/16] introduce base class for shampoo Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 67 +++++++++++++++++++++----- 1 file changed, 54 insertions(+), 13 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index 2d17701..ae0a940 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -23,12 +23,14 @@ from torch.optim.optimizer import ParamsT from emerging_optimizers import mixin as opt_mixin +from emerging_optimizers import registry from emerging_optimizers.shampoo import shampoo_base from emerging_optimizers.utils import eig as eig_utils __all__ = [ "Shampoo", + "ShampooBase", "ShampooPreconditioner", ] @@ -102,6 +104,9 @@ def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: grad: Gradient of the parameter on the first step. shampoo_beta: EMA coefficient for the Kronecker factor update. """ + # Changing initial kronecker factors to be epsilon along the diagonal to match the paper + self.kronecker_factor_pair.L += self.eps * torch.eye(grad.shape[-2], device=grad.device) + self.kronecker_factor_pair.R += self.eps * torch.eye(grad.shape[-1], device=grad.device) self.update_kronecker_factors(grad, shampoo_beta) def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: @@ -123,7 +128,7 @@ def step(self, grad: torch.Tensor, shampoo_beta: float) -> None: """ self.update_kronecker_factors(grad, shampoo_beta) - def _get_inverse_root(self, kronecker_factor: torch.Tensor) -> torch.Tensor: + def _get_root_inverse(self, kronecker_factor: torch.Tensor) -> torch.Tensor: """Computes ``kronecker_factor^(-1/p_inv_root)`` from its eigendecomposition. Args: @@ -133,8 +138,8 @@ def _get_inverse_root(self, kronecker_factor: torch.Tensor) -> torch.Tensor: The inverse root of the factor. """ eigvals, eigvecs = eig_utils.eigh_with_fallback(kronecker_factor) - # Kronecker factors are symmetric. Eigh can sometime return negative values for numerical 0 therefore the abs. - return (eigvecs * eigvals.abs().clamp_min(self.eps) ** (-1.0 / self.p_inv_root)) @ eigvecs.mT + # Eigh can sometime return negative values for numerical 0, which clamp to eps will also get rid off + return (eigvecs * eigvals.clamp_min(self.eps) ** (-1.0 / self.p_inv_root)) @ eigvecs.mT def precondition(self, x: torch.Tensor) -> torch.Tensor: """Applies both inverse roots to a matrix in the parameter basis. @@ -146,18 +151,22 @@ def precondition(self, x: torch.Tensor) -> torch.Tensor: The preconditioned matrix, in the parameter basis. """ inverse_root_pair = shampoo_base.TensorPair( - self._get_inverse_root(self.kronecker_factor_pair.L), - self._get_inverse_root(self.kronecker_factor_pair.R), + self._get_root_inverse(self.kronecker_factor_pair.L), + self._get_root_inverse(self.kronecker_factor_pair.R), ) return inverse_root_pair.L @ x @ inverse_root_pair.R -class Shampoo(optim.Optimizer, opt_mixin.WeightDecayMixin): - """Shampoo with inverse roots rebuilt by eigendecomposition on every step. +class ShampooBase(optim.Optimizer, opt_mixin.WeightDecayMixin): + """Canonical Shampoo step loop, shared by the Shampoo-family optimizers. - The update is EMA momentum in the parameter basis, preconditioned on both sides by the inverse roots - of the Kronecker factors. + :meth:`step` is the whole algorithm: update the preconditioner from the gradient, run an inner scalar + optimizer in the parameter basis, and precondition its update on both sides. Subclasses customize the + two pieces that vary between Shampoo variants and leave the loop alone: + + - :attr:`PreconditionerCls` -- how the Kronecker factors and their inverse roots are maintained. + - :meth:`_scalar_update` -- which scalar optimizer produces the update being preconditioned. Args: params: Iterable of 2D CUDA parameters to optimize or dicts defining parameter groups. @@ -172,7 +181,7 @@ class Shampoo(optim.Optimizer, opt_mixin.WeightDecayMixin): allocated by :meth:`_init_group`. Subclasses set it to change how the factors are maintained. """ - PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] = ShampooPreconditioner + PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] def __init__( self, @@ -198,7 +207,6 @@ def __init__( } super().__init__(params, defaults) - @torch.compile def _scalar_update( self, grad: torch.Tensor, @@ -208,6 +216,8 @@ def _scalar_update( ) -> torch.Tensor: """Applies the inner scalar optimizer to the gradient, in the parameter basis. + Override this to run a different scalar update ahead of the preconditioner. + Args: grad: Gradient of the parameter. exp_avg: Momentum buffer, updated in place. @@ -215,9 +225,11 @@ def _scalar_update( Returns: The scalar update, in the parameter basis. + + Raises: + NotImplementedError: Always; subclasses must provide the inner update. """ - exp_avg.lerp_(grad, 1 - momentum) - return exp_avg + raise NotImplementedError @torch.no_grad() # type: ignore[misc] def _init_group( @@ -306,3 +318,32 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: state["step"] += 1 return None + + +@registry.register_optimizer("shampoo") +class Shampoo(ShampooBase): + """Shampoo with EMA momentum as the inner scalar optimizer.""" + + PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] = ShampooPreconditioner + + @torch.compile + @override + def _scalar_update( + self, + grad: torch.Tensor, + exp_avg: torch.Tensor, + *, + momentum: float, + ) -> torch.Tensor: + """Applies EMA momentum to the gradient, in the parameter basis. + + Args: + grad: Gradient of the parameter. + exp_avg: Momentum buffer, updated in place. + momentum: Momentum EMA coefficient. + + Returns: + The momentum update, in the parameter basis. + """ + exp_avg.lerp_(grad, 1 - momentum) + return exp_avg From 72321f939c773f771162242c1f1840c6dc877c4a Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 13 Aug 2026 16:06:22 -0700 Subject: [PATCH 04/16] save work Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index ae0a940..d7dd6b1 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -48,7 +48,7 @@ def __init__( self, state: dict, p_inv_root: float, - eps: float = 1e-8, + eps: float, ) -> None: self.kronecker_factor_pair = shampoo_base.TensorPair(state["L"], state["R"]) self.p_inv_root = p_inv_root @@ -173,6 +173,7 @@ class ShampooBase(optim.Optimizer, opt_mixin.WeightDecayMixin): lr: Learning rate. momentum: Momentum EMA coefficient. shampoo_beta: Kronecker factor EMA coefficient. + eps: Numerical epsilon weight_decay: Decoupled weight decay coefficient. p_inv_root: Inverse root order applied to each Kronecker factor. @@ -189,15 +190,19 @@ def __init__( lr: float, momentum: float = 0.9, shampoo_beta: float = 0.95, + eps: float = 1e-8, weight_decay: float = 0.01, *, p_inv_root: float = 4, ) -> None: + self.eps = eps self.weight_decay_method = "decoupled" self.p_inv_root = p_inv_root if lr < 0.0: raise ValueError(f"Invalid learning rate: {lr}") + if p_inv_root <= 0 or round(p_inv_root) != p_inv_root: + raise ValueError(f"p_inv_root must be positive integer, got {p_inv_root}") defaults = { "lr": lr, @@ -244,7 +249,7 @@ def _init_group( skip_non_grad_params: Whether to skip parameters with no gradients. Raises: - TypeError: If the parameter is not a 2D CUDA tensor. + TypeError: If the parameter is not a 2D tensor. """ for p in group["params"]: if skip_non_grad_params and p.grad is None: @@ -252,8 +257,6 @@ def _init_group( if p.dim() != 2: raise TypeError(f"{type(self).__name__} is only supported for 2D tensors") - if not p.is_cuda: - raise TypeError(f"{type(self).__name__} only supports CUDA tensors") state = self.state[p] @@ -279,7 +282,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: Args: closure: Unsupported; must be ``None``. - Raises: + Raises:f ValueError: If ``closure`` is not ``None``. """ if closure is not None: @@ -296,7 +299,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: grad = p.grad.to(torch.float32) state = self.state[p] - preconditioner = self.PreconditionerCls(state, self.p_inv_root) + preconditioner = self.PreconditionerCls(state, self.p_inv_root, self.eps) scalar_update = self._scalar_update(grad, state["exp_avg"], momentum=group["momentum"]) From 5462eb40d000e2d6e881448afcb56a18d28befd6 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 14 Aug 2026 10:07:01 -0700 Subject: [PATCH 05/16] add test for vanilla shampoo Signed-off-by: Hao Wu --- tests/test_shampoo.py | 359 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 359 insertions(+) create mode 100644 tests/test_shampoo.py diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py new file mode 100644 index 0000000..9ed04f5 --- /dev/null +++ b/tests/test_shampoo.py @@ -0,0 +1,359 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import override + +import torch +from _comparison import assert_equal +from absl import flags, logging +from absl.testing import absltest, parameterized + +from emerging_optimizers import utils +from emerging_optimizers.legacy_soap import soap +from emerging_optimizers.shampoo.shampoo import Shampoo, ShampooBase, ShampooPreconditioner + + +flags.DEFINE_enum("device", "cpu", ["cpu", "cuda"], "Device to run tests on") +flags.DEFINE_integer("seed", None, "Random seed for reproducible tests") +FLAGS = flags.FLAGS + + +def setUpModule() -> None: + if FLAGS.seed is not None: + logging.info("Setting random seed to %d", FLAGS.seed) + torch.manual_seed(FLAGS.seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(FLAGS.seed) + + +def _root_inverse_reference(a: torch.Tensor, p_inv_root: float, eps: float) -> torch.Tensor: + u, s, vh = torch.linalg.svd(a) + return (u * s.clamp_min(eps) ** (-1.0 / p_inv_root)) @ vh + + +def gen_signed_permutation(m: int): + signs = torch.randint(0, 2, (m,), dtype=torch.float32) * 2 - 1 + Q = torch.zeros(m, m) + Q[torch.randperm(m), torch.arange(m)] = signs + + return Q + + +class ShampooPreconditionerTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + cls.device = torch.device(FLAGS.device) + + @parameterized.parameters((8, 16), (16, 8), (13, 15)) + def test_init_state_layout(self, m: int, n: int) -> None: + state = ShampooPreconditioner.init_state((m, n), self.device) + + expected_shapes = {"L": (m, m), "R": (n, n)} + self.assertCountEqual(state, expected_shapes) + for key, shape in expected_shapes.items(): + self.assertEqual(state[key].shape, shape, msg=key) + self.assertEqual(state[key].dtype, torch.float32, msg=key) + self.assertEqual(state[key].device.type, self.device.type, msg=key) + assert_equal(state[key], torch.zeros(shape, device=self.device)) + + def test_init_state_rejects_non_2d(self) -> None: + with self.assertRaisesRegex(TypeError, "only supported for 2D"): + ShampooPreconditioner.init_state((2, 3, 4), self.device) + + @parameterized.parameters((8, 16), (16, 8), (13, 15)) + def test_rebind_state_binds_current_tensors(self, m: int, n: int) -> None: + state = ShampooPreconditioner.init_state((m, n), self.device) + preconditioner = ShampooPreconditioner(state, p_inv_root=4, eps=1e-8) + preconditioner.step(torch.randn(m, n, device=self.device), 0.95) + preconditioner.rebind_state(state) + + self.assertIs(state["L"], preconditioner.kronecker_factor_pair.L) + self.assertIs(state["R"], preconditioner.kronecker_factor_pair.R) + + def test_rebind_state_missing_key_raises(self) -> None: + state = ShampooPreconditioner.init_state((4, 4), self.device) + preconditioner = ShampooPreconditioner(state, p_inv_root=4, eps=1e-8) + del state["R"] + + with self.assertRaisesRegex(KeyError, "missing keys"): + preconditioner.rebind_state(state) + + @parameterized.parameters((8, 16), (16, 8), (13, 15)) + def test_init_step_seeds_eps_identity(self, m: int, n: int) -> None: + eps = 0.5 + preconditioner = ShampooPreconditioner( + ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps + ) + grad = torch.zeros(m, n, device=self.device) + + preconditioner.init_step(grad, shampoo_beta=1) + + assert_equal(preconditioner.kronecker_factor_pair.L, torch.eye(m, device=self.device) * eps) + assert_equal(preconditioner.kronecker_factor_pair.R, torch.eye(n, device=self.device) * eps) + + @parameterized.product(shape=[(8, 16), (16, 8), (13, 15)], shampoo_beta=[0.5, 0.95]) + def test_update_kronecker_factors_matches_legacy(self, shape: tuple[int, int], shampoo_beta: float) -> None: + m, n = shape + preconditioner = ShampooPreconditioner( + ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=1e-8 + ) + preconditioner.init_step(torch.randn(m, n, device=self.device), shampoo_beta) + + reference_factors = [ + preconditioner.kronecker_factor_pair.L.clone(), + preconditioner.kronecker_factor_pair.R.clone(), + ] + grad = torch.randn(m, n, device=self.device) + soap.update_kronecker_factors(reference_factors, grad, shampoo_beta) + preconditioner.update_kronecker_factors(grad, shampoo_beta) + + assert_equal(preconditioner.kronecker_factor_pair.L, reference_factors[0]) + assert_equal(preconditioner.kronecker_factor_pair.R, reference_factors[1]) + + def test_step_equals_update_kronecker_factors(self) -> None: + eps = 1e-8 + m, n, shampoo_beta = 6, 4, 0.9 + grad = torch.randn(m, n, device=self.device) + + stepped = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps) + updated = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps) + stepped.step(grad, shampoo_beta) + updated.update_kronecker_factors(grad, shampoo_beta) + + assert_equal(stepped.kronecker_factor_pair.L, updated.kronecker_factor_pair.L) + assert_equal(stepped.kronecker_factor_pair.R, updated.kronecker_factor_pair.R) + + @parameterized.product(m=[4, 9, 16], p_inv_root=[2, 4]) + def test_get_root_inverse_close_to_svd_reference(self, m: int, p_inv_root: int) -> None: + x = 2 ** torch.randint(-3, 2, (m, m), device=self.device, dtype=torch.float) + factor = x @ x.T + 0.125 * torch.eye(m, device=self.device) + preconditioner = ShampooPreconditioner( + ShampooPreconditioner.init_state((m, m), self.device), p_inv_root=p_inv_root, eps=0 + ) + + with utils.fp32_matmul_precision("highest"): + root_inverse = preconditioner._get_root_inverse(factor) + + torch.testing.assert_close( + root_inverse, + _root_inverse_reference(factor, p_inv_root, 0), + atol=1e-3, + rtol=1e-3, + ) + + @parameterized.parameters((6, 4), (4, 6), (5, 5)) + def test_precondition_identity_factors_is_noop(self, m: int, n: int) -> None: + preconditioner = ShampooPreconditioner( + {"L": torch.eye(m, device=self.device), "R": torch.eye(n, device=self.device)}, p_inv_root=4, eps=1e-8 + ) + x = torch.randn(m, n, device=self.device) + + with utils.fp32_matmul_precision("highest"): + preconditioned = preconditioner.precondition(x) + + assert_equal(preconditioned, x) + + @parameterized.parameters(6, 16, 33) + def test_precondition_matches_inverse_of_known_spectrum(self, m: int) -> None: + """Test designed to have exact match. + + Kronecker factors are created by integer eigen values and signed permutation eigven vectors. + """ + p = gen_signed_permutation(m).to(self.device) + eigvals = 2 ** torch.randint(-5, 0, (m,), dtype=torch.float32, device=self.device) + A = p * eigvals @ p.mT + + init_kronecker_factors = { + "L": A, + "R": A.clone(), + } + inv_root_kwargs = { + "p_inv_root": 2, + "eps": 0, + } + preconditioner = ShampooPreconditioner(init_kronecker_factors, **inv_root_kwargs) + + scale = 7 + x = torch.eye(m, device=self.device, dtype=torch.float32) * scale + with utils.fp32_matmul_precision("highest"): + preconditioned = preconditioner.precondition(x).round() + + expected = (p * (eigvals**-1) @ p.mT * scale).round() + + assert_equal(preconditioned, expected) + + @parameterized.parameters((8, 3), (3, 8)) + def test_precondition_4steps_smoke(self, m: int, n: int) -> None: + shampoo_beta = 0.95 + preconditioner = ShampooPreconditioner( + ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=1e-8 + ) + preconditioner.init_step(torch.randn(m, n, device=self.device), shampoo_beta) + for _ in range(4): + preconditioner.step(torch.randn(m, n, device=self.device), shampoo_beta) + + preconditioned = preconditioner.precondition(torch.randn(m, n, device=self.device)) + + self.assertEqual(preconditioned.shape, (m, n)) + + +class _BypassPreconditioner: + def __init__(self, state: dict, p_inv_root: float, eps: float) -> None: + self.p_inv_root = p_inv_root + self.eps = eps + + @staticmethod + def init_state(shape: tuple[int, ...], device: torch.device) -> dict[str, torch.Tensor]: + return {} + + def rebind_state(self, state: dict) -> None: + pass + + def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: + pass + + def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: + pass + + def step(self, grad: torch.Tensor, shampoo_beta: float) -> None: + pass + + def precondition(self, x: torch.Tensor) -> torch.Tensor: + return x + + +class _SgdShampoo(ShampooBase): + """A fake shampoo that bypass preconditioning for testing base class.""" + + PreconditionerCls = _BypassPreconditioner + + @override + def _scalar_update( + self, + grad: torch.Tensor, + exp_avg: torch.Tensor, + *, + momentum: float, + ) -> torch.Tensor: + exp_avg.mul_(momentum).add_(grad) + return exp_avg + + +class ShampooBaseTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + cls.device = torch.device(FLAGS.device) + + def test_step_smoke(self) -> None: + p = torch.nn.Parameter(torch.randn(4, 4, device=self.device)) + optimizer = _SgdShampoo([p], lr=1e-3) + + p.grad = torch.randn_like(p) + + optimizer.step() + + @parameterized.product(lr=[0.25, 0.125], momentum=[0.0, 75], weight_decay=[0.125, 0.05]) + def test_step_3steps_close_to_sgd(self, lr: float, momentum: float, weight_decay: float) -> None: + p = torch.nn.Parameter(torch.randn(4, 3, device=self.device)) + expected = p.detach().clone() + sgd_buffer = torch.zeros_like(expected) + optimizer = _SgdShampoo([p], lr=lr, momentum=momentum, weight_decay=weight_decay) + + for _ in range(3): + grad = torch.randn_like(p) + p.grad = grad.clone() + optimizer.step() + sgd_buffer = momentum * sgd_buffer + grad + expected.mul_(1 - lr * weight_decay).add_(sgd_buffer, alpha=-lr) + + torch.testing.assert_close( + p.detach(), + expected, + atol=1e-5, + rtol=1e-5, + ) + + def test_rejects_non_2d(self) -> None: + p = torch.nn.Parameter(torch.randn(2, 3, 4, device=self.device)) + p.grad = torch.randn_like(p) + optimizer = _SgdShampoo([p], lr=1e-3) + + with self.assertRaisesRegex(TypeError, "only supported for 2D"): + optimizer.step() + + def test_negative_lr_raises(self) -> None: + with self.assertRaisesRegex(ValueError, "Invalid learning rate"): + _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], lr=-1.0) + + +class ShampooTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + cls.device = torch.device(FLAGS.device) + + @parameterized.parameters((8, 5), (5, 8), (16, 16)) + def test_3steps_smoke(self, m: int, n: int) -> None: + p = torch.nn.Parameter(torch.randn(m, n, device=self.device)) + initial = p.detach().clone() + + optimizer = Shampoo([p], lr=1e-2) + for _ in range(3): + p.grad = torch.randn_like(p) + optimizer.step() + + self.assertTrue(torch.isfinite(p).all()) + self.assertFalse(torch.equal(p.detach(), initial)) + state = optimizer.state[p] + self.assertEqual(state["step"], 3) + self.assertCountEqual(state, {"step", "exp_avg", "L", "R"}) + + def test_zero_grad_applies_only_weight_decay(self) -> None: + lr, weight_decay = 0.1, 0.05 + p = torch.nn.Parameter(torch.randn(5, 5, device=self.device)) + initial = p.detach().clone() + optimizer = Shampoo([p], lr=lr, weight_decay=weight_decay) + + p.grad = torch.zeros_like(p) + optimizer.step() + + torch.testing.assert_close( + p.detach(), + initial * (1 - lr * weight_decay), + atol=1e-6, + rtol=1e-6, + msg=lambda default: f"A zero gradient should leave only decoupled weight decay\n\n{default}", + ) + + @parameterized.parameters(0.0, 0.5, 0.75) + def test_scalar_update_close_to_sgd(self, momentum: float) -> None: + p = torch.nn.Parameter(torch.randn(6, 6, device=self.device)) + optimizer = Shampoo([p], lr=0.125, momentum=momentum) + exp_avg = torch.zeros_like(p) + sgd_buffer = torch.zeros_like(p) + + for _ in range(3): + grad = torch.randn_like(p) + scalar_update = optimizer._scalar_update(grad, exp_avg, momentum=momentum) + sgd_buffer.mul_(momentum).add_(grad, alpha=1 - momentum) + + torch.testing.assert_close( + scalar_update, + sgd_buffer, + atol=1e-5, + rtol=1e-5, + ) + + +if __name__ == "__main__": + absltest.main() From 6993ab4e0c6510ad5f01f2e0a9a0878ba7584aec Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 14 Aug 2026 10:09:59 -0700 Subject: [PATCH 06/16] rename module Signed-off-by: Hao Wu --- .../{shampoo_base.py => precond_base.py} | 0 emerging_optimizers/shampoo/shampoo.py | 10 +++---- emerging_optimizers/shampoo/soap_v3.py | 30 +++++++++---------- 3 files changed, 20 insertions(+), 20 deletions(-) rename emerging_optimizers/shampoo/{shampoo_base.py => precond_base.py} (100%) diff --git a/emerging_optimizers/shampoo/shampoo_base.py b/emerging_optimizers/shampoo/precond_base.py similarity index 100% rename from emerging_optimizers/shampoo/shampoo_base.py rename to emerging_optimizers/shampoo/precond_base.py diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index d7dd6b1..2e90541 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -24,7 +24,7 @@ from emerging_optimizers import mixin as opt_mixin from emerging_optimizers import registry -from emerging_optimizers.shampoo import shampoo_base +from emerging_optimizers.shampoo import precond_base from emerging_optimizers.utils import eig as eig_utils @@ -50,7 +50,7 @@ def __init__( p_inv_root: float, eps: float, ) -> None: - self.kronecker_factor_pair = shampoo_base.TensorPair(state["L"], state["R"]) + self.kronecker_factor_pair = precond_base.TensorPair(state["L"], state["R"]) self.p_inv_root = p_inv_root self.eps = eps @@ -150,7 +150,7 @@ def precondition(self, x: torch.Tensor) -> torch.Tensor: Returns: The preconditioned matrix, in the parameter basis. """ - inverse_root_pair = shampoo_base.TensorPair( + inverse_root_pair = precond_base.TensorPair( self._get_root_inverse(self.kronecker_factor_pair.L), self._get_root_inverse(self.kronecker_factor_pair.R), ) @@ -182,7 +182,7 @@ class ShampooBase(optim.Optimizer, opt_mixin.WeightDecayMixin): allocated by :meth:`_init_group`. Subclasses set it to change how the factors are maintained. """ - PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] + PreconditionerCls: ClassVar[type[precond_base.ShampooPreconditionerProtocol]] def __init__( self, @@ -327,7 +327,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: class Shampoo(ShampooBase): """Shampoo with EMA momentum as the inner scalar optimizer.""" - PreconditionerCls: ClassVar[type[shampoo_base.ShampooPreconditionerProtocol]] = ShampooPreconditioner + PreconditionerCls: ClassVar[type[precond_base.ShampooPreconditionerProtocol]] = ShampooPreconditioner @torch.compile @override diff --git a/emerging_optimizers/shampoo/soap_v3.py b/emerging_optimizers/shampoo/soap_v3.py index c52b1c3..07ea12a 100644 --- a/emerging_optimizers/shampoo/soap_v3.py +++ b/emerging_optimizers/shampoo/soap_v3.py @@ -27,7 +27,7 @@ from emerging_optimizers import registry, utils from emerging_optimizers.legacy_soap import soap from emerging_optimizers.scalar_optimizers import update_functions -from emerging_optimizers.shampoo import shampoo_base +from emerging_optimizers.shampoo import precond_base from emerging_optimizers.utils import eig as eig_utils @@ -54,9 +54,9 @@ def __init__( state: dict, eps: float, ) -> None: - self.kronecker_factor_pair = shampoo_base.TensorPair(state["L"], state["R"]) - self.eigenbasis_pair = shampoo_base.TensorPair(state["Q_L"], state["Q_R"]) - self.eigvals_pair = shampoo_base.TensorPair(state["eigvals_L"], state["eigvals_R"]) + self.kronecker_factor_pair = precond_base.TensorPair(state["L"], state["R"]) + self.eigenbasis_pair = precond_base.TensorPair(state["Q_L"], state["Q_R"]) + self.eigvals_pair = precond_base.TensorPair(state["eigvals_L"], state["eigvals_R"]) self.exp_avg, self.exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] self.eps = eps @@ -124,8 +124,8 @@ def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: self.update_kronecker_factors(grad, shampoo_beta) eigvals_L, Q_L = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.L) eigvals_R, Q_R = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.R) - self.eigenbasis_pair = shampoo_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = shampoo_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) + self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: """Accumulates the gradient into the kronecker factors with the KL-Shampoo correction. @@ -169,8 +169,8 @@ def step( eigvals_R, Q_R = eig_utils.orthogonal_iteration( self.kronecker_factor_pair.R, self.eigenbasis_pair.R, power_iter_steps=1 ) - self.eigenbasis_pair = shampoo_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = shampoo_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) + self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) # Project exp_avg to the new eigenbasis using the updated eigenbases self.exp_avg = self.project_in(exp_avg) @@ -223,8 +223,8 @@ def step( # Rebuild the eigen bases from the factors rather than refining the previous ones eigvals_L, Q_L = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.L) eigvals_R, Q_R = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.R) - self.eigenbasis_pair = shampoo_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = shampoo_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) + self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) # Project exp_avg to the new eigenbasis using the updated eigenbases self.exp_avg = self.project_in(exp_avg) @@ -259,12 +259,12 @@ class SoapBase(optim.Optimizer, opt_mixin.WeightDecayMixin): Attributes: PreconditionerCls: Preconditioner used for every parameter. Subclasses set it to change how the covariance factors and eigenbases are maintained; it must satisfy - :class:`~emerging_optimizers.shampoo.shampoo_base.SoapPreconditionerProtocol`. It is + :class:`~emerging_optimizers.shampoo.precond_base.SoapPreconditionerProtocol`. It is also what :meth:`_init_group` allocates state from, so a subclass that swaps it gets that preconditioner's state layout. """ - PreconditionerCls: ClassVar[type[shampoo_base.SoapPreconditionerProtocol]] + PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] def __init__( self, @@ -445,7 +445,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: class KlSoapV3(SoapBase): """Implements a variant of KLSOAP algorithm.""" - PreconditionerCls: ClassVar[type[shampoo_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner + PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner @override def _scalar_update( @@ -485,14 +485,14 @@ def _scalar_update( class ReklsV3(KlSoapV3): """Realtime Eigen KL-Shampoo""" - PreconditionerCls: ClassVar[type[shampoo_base.SoapPreconditionerProtocol]] = ReklsPreconditioner + PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = ReklsPreconditioner @registry.register_optimizer("kl_m_soap") class KlMSoap(SoapBase): """SOAP with the KL-Shampoo kronecker factor update and MAdam as the inner scalar optimizer.""" - PreconditionerCls: ClassVar[type[shampoo_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner + PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner @override def _scalar_update( From 3e38538717ce4fe0d985ae8eb99b5f09037645e8 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 14 Aug 2026 10:46:42 -0700 Subject: [PATCH 07/16] fix error type Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/soap_v3.py | 2 +- tests/test_soap_v3.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/emerging_optimizers/shampoo/soap_v3.py b/emerging_optimizers/shampoo/soap_v3.py index 07ea12a..dffb891 100644 --- a/emerging_optimizers/shampoo/soap_v3.py +++ b/emerging_optimizers/shampoo/soap_v3.py @@ -75,7 +75,7 @@ def init_state( The state entries owned by this preconditioner, keyed as :meth:`rebind_state` expects them. Raises: - ValueError: If ``shape`` is not 2D. + TypeError: If ``shape`` is not 2D. """ if len(shape) != 2: raise TypeError(f"KlSoapPreconditioner is only supported for 2D tensors, got shape {tuple(shape)}") diff --git a/tests/test_soap_v3.py b/tests/test_soap_v3.py index cdd7974..159c459 100644 --- a/tests/test_soap_v3.py +++ b/tests/test_soap_v3.py @@ -60,7 +60,7 @@ def test_init_state_layout(self, m: int, n: int) -> None: assert_equal(state["Q_R"], torch.eye(n, device=FLAGS.device)) def test_init_state_rejects_non_2d(self) -> None: - with self.assertRaisesRegex(ValueError, "only supported for 2D"): + with self.assertRaisesRegex(TypeError, "only supported for 2D"): KlSoapPreconditioner.init_state((2, 3, 4), torch.device(FLAGS.device)) @parameterized.parameters((8, 16), (16, 8), (12, 12)) From 20b78d057e98e29ba92a66807171bb19ed519a97 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 14 Aug 2026 11:18:20 -0700 Subject: [PATCH 08/16] improve test coverage Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/precond_base.py | 4 ---- tests/test_shampoo.py | 17 +++++++++++++++++ 2 files changed, 17 insertions(+), 4 deletions(-) diff --git a/emerging_optimizers/shampoo/precond_base.py b/emerging_optimizers/shampoo/precond_base.py index 5f9e675..77dbcb1 100644 --- a/emerging_optimizers/shampoo/precond_base.py +++ b/emerging_optimizers/shampoo/precond_base.py @@ -37,10 +37,6 @@ def __iter__(self) -> Iterator[torch.Tensor]: """Iterates over the pair as ``L`` then ``R``.""" return iter((self.L, self.R)) - def __getitem__(self, index: int) -> torch.Tensor: - """Indexes the pair as ``0`` for ``L`` and ``1`` for ``R``.""" - return (self.L, self.R)[index] - class _PreconditionerProtocol(Protocol): """Interface every preconditioner in the family must provide, for one parameter. diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 9ed04f5..1ebf4f9 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -318,6 +318,23 @@ def test_3steps_smoke(self, m: int, n: int) -> None: self.assertEqual(state["step"], 3) self.assertCountEqual(state, {"step", "exp_avg", "L", "R"}) + @parameterized.parameters(True, False) + def test_init_group_skip_non_grad_params(self, skip_non_grad_params: bool) -> None: + with_grad = torch.nn.Parameter(torch.randn(4, 3, device=self.device)) + without_grad = torch.nn.Parameter(torch.randn(5, 2, device=self.device)) + with_grad.grad = torch.randn_like(with_grad) + optimizer = Shampoo([with_grad, without_grad], lr=1e-3) + + optimizer._init_group(optimizer.param_groups[0], skip_non_grad_params=skip_non_grad_params) + + self.assertCountEqual(optimizer.state[with_grad], {"step", "exp_avg", "L", "R"}) + if skip_non_grad_params: + self.assertEmpty(optimizer.state[without_grad]) + else: + self.assertCountEqual(optimizer.state[without_grad], {"step", "exp_avg", "L", "R"}) + self.assertEqual(optimizer.state[without_grad]["L"].shape, (5, 5)) + self.assertEqual(optimizer.state[without_grad]["R"].shape, (2, 2)) + def test_zero_grad_applies_only_weight_decay(self) -> None: lr, weight_decay = 0.1, 0.05 p = torch.nn.Parameter(torch.randn(5, 5, device=self.device)) From edb7040fa9f75c64cebc410abff9bfd087485f24 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 17 Aug 2026 14:19:02 -0700 Subject: [PATCH 09/16] update import order Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/__init__.py | 3 +++ emerging_optimizers/shampoo/soap_v3.py | 30 ++++++++++++------------- 2 files changed, 18 insertions(+), 15 deletions(-) diff --git a/emerging_optimizers/shampoo/__init__.py b/emerging_optimizers/shampoo/__init__.py index 4670798..02a61d0 100644 --- a/emerging_optimizers/shampoo/__init__.py +++ b/emerging_optimizers/shampoo/__init__.py @@ -12,3 +12,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. + +from emerging_optimizers.shampoo.shampoo import Shampoo +from emerging_optimizers.shampoo.soap_v3 import KlMSoap, KlSoapV3, ReklsV3 diff --git a/emerging_optimizers/shampoo/soap_v3.py b/emerging_optimizers/shampoo/soap_v3.py index dffb891..4d13658 100644 --- a/emerging_optimizers/shampoo/soap_v3.py +++ b/emerging_optimizers/shampoo/soap_v3.py @@ -27,7 +27,7 @@ from emerging_optimizers import registry, utils from emerging_optimizers.legacy_soap import soap from emerging_optimizers.scalar_optimizers import update_functions -from emerging_optimizers.shampoo import precond_base +from emerging_optimizers.shampoo.precond_base import SoapPreconditionerProtocol, TensorPair from emerging_optimizers.utils import eig as eig_utils @@ -54,9 +54,9 @@ def __init__( state: dict, eps: float, ) -> None: - self.kronecker_factor_pair = precond_base.TensorPair(state["L"], state["R"]) - self.eigenbasis_pair = precond_base.TensorPair(state["Q_L"], state["Q_R"]) - self.eigvals_pair = precond_base.TensorPair(state["eigvals_L"], state["eigvals_R"]) + self.kronecker_factor_pair = TensorPair(state["L"], state["R"]) + self.eigenbasis_pair = TensorPair(state["Q_L"], state["Q_R"]) + self.eigvals_pair = TensorPair(state["eigvals_L"], state["eigvals_R"]) self.exp_avg, self.exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] self.eps = eps @@ -124,8 +124,8 @@ def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: self.update_kronecker_factors(grad, shampoo_beta) eigvals_L, Q_L = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.L) eigvals_R, Q_R = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.R) - self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = TensorPair(Q_L, Q_R) + self.eigvals_pair = TensorPair(eigvals_L, eigvals_R) def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: """Accumulates the gradient into the kronecker factors with the KL-Shampoo correction. @@ -169,8 +169,8 @@ def step( eigvals_R, Q_R = eig_utils.orthogonal_iteration( self.kronecker_factor_pair.R, self.eigenbasis_pair.R, power_iter_steps=1 ) - self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = TensorPair(Q_L, Q_R) + self.eigvals_pair = TensorPair(eigvals_L, eigvals_R) # Project exp_avg to the new eigenbasis using the updated eigenbases self.exp_avg = self.project_in(exp_avg) @@ -223,8 +223,8 @@ def step( # Rebuild the eigen bases from the factors rather than refining the previous ones eigvals_L, Q_L = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.L) eigvals_R, Q_R = eig_utils.eigh_with_fallback(self.kronecker_factor_pair.R) - self.eigenbasis_pair = precond_base.TensorPair(Q_L, Q_R) - self.eigvals_pair = precond_base.TensorPair(eigvals_L, eigvals_R) + self.eigenbasis_pair = TensorPair(Q_L, Q_R) + self.eigvals_pair = TensorPair(eigvals_L, eigvals_R) # Project exp_avg to the new eigenbasis using the updated eigenbases self.exp_avg = self.project_in(exp_avg) @@ -259,12 +259,12 @@ class SoapBase(optim.Optimizer, opt_mixin.WeightDecayMixin): Attributes: PreconditionerCls: Preconditioner used for every parameter. Subclasses set it to change how the covariance factors and eigenbases are maintained; it must satisfy - :class:`~emerging_optimizers.shampoo.precond_base.SoapPreconditionerProtocol`. It is + :class:`~emerging_optimizers.shampoo.SoapPreconditionerProtocol`. It is also what :meth:`_init_group` allocates state from, so a subclass that swaps it gets that preconditioner's state layout. """ - PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] + PreconditionerCls: ClassVar[type[SoapPreconditionerProtocol]] def __init__( self, @@ -445,7 +445,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: class KlSoapV3(SoapBase): """Implements a variant of KLSOAP algorithm.""" - PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner + PreconditionerCls: ClassVar[type[SoapPreconditionerProtocol]] = KlSoapPreconditioner @override def _scalar_update( @@ -485,14 +485,14 @@ def _scalar_update( class ReklsV3(KlSoapV3): """Realtime Eigen KL-Shampoo""" - PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = ReklsPreconditioner + PreconditionerCls: ClassVar[type[SoapPreconditionerProtocol]] = ReklsPreconditioner @registry.register_optimizer("kl_m_soap") class KlMSoap(SoapBase): """SOAP with the KL-Shampoo kronecker factor update and MAdam as the inner scalar optimizer.""" - PreconditionerCls: ClassVar[type[precond_base.SoapPreconditionerProtocol]] = KlSoapPreconditioner + PreconditionerCls: ClassVar[type[SoapPreconditionerProtocol]] = KlSoapPreconditioner @override def _scalar_update( From 91d5c093e51f0ed1246d21ae3cdc21cc6fdf2460 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 17 Aug 2026 14:36:49 -0700 Subject: [PATCH 10/16] add tikhonov and bias correction Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index 2e90541..1189727 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -138,8 +138,14 @@ def _get_root_inverse(self, kronecker_factor: torch.Tensor) -> torch.Tensor: The inverse root of the factor. """ eigvals, eigvecs = eig_utils.eigh_with_fallback(kronecker_factor) - # Eigh can sometime return negative values for numerical 0, which clamp to eps will also get rid off - return (eigvecs * eigvals.clamp_min(self.eps) ** (-1.0 / self.p_inv_root)) @ eigvecs.mT + + # Eigh can sometime return negative values for numerical 0; clamping to 0 removes them + eigvals = eigvals.clamp_min(0) + + # Tikhonov regularization + exp = 1.0 / self.p_inv_root + inv_root_scale = eigvals**exp / (eigvals ** (2 * exp) + self.eps ** (2 * exp)) + return (eigvecs * inv_root_scale) @ eigvecs.mT def precondition(self, x: torch.Tensor) -> torch.Tensor: """Applies both inverse roots to a matrix in the parameter basis. @@ -303,6 +309,11 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: scalar_update = self._scalar_update(grad, state["exp_avg"], momentum=group["momentum"]) + # bias correction on shampoo beta + curr_iter_1_based = state["step"] + 1 + shampoo_beta = group["shampoo_beta"] + shampoo_beta = 1 - (1 - shampoo_beta) / (1 - shampoo_beta**curr_iter_1_based) + if state["step"] == 0: preconditioner.init_step(grad, group["shampoo_beta"]) else: From 5e2ed4e99b24b65561b3ac1ddc391cc7890493ac Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 18 Aug 2026 12:29:08 -0700 Subject: [PATCH 11/16] fix tests Signed-off-by: Hao Wu --- tests/test_shampoo.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 1ebf4f9..3378ad3 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -155,7 +155,7 @@ def test_get_root_inverse_close_to_svd_reference(self, m: int, p_inv_root: int) @parameterized.parameters((6, 4), (4, 6), (5, 5)) def test_precondition_identity_factors_is_noop(self, m: int, n: int) -> None: preconditioner = ShampooPreconditioner( - {"L": torch.eye(m, device=self.device), "R": torch.eye(n, device=self.device)}, p_inv_root=4, eps=1e-8 + {"L": torch.eye(m, device=self.device), "R": torch.eye(n, device=self.device)}, p_inv_root=4, eps=0 ) x = torch.randn(m, n, device=self.device) From b6140c5deda7ecd8a07a304d53e30bcb9077d78d Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 18 Aug 2026 12:53:53 -0700 Subject: [PATCH 12/16] add test for shampoo beta bias correction Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 4 ++-- tests/test_shampoo.py | 25 ++++++++++++++++++++++--- 2 files changed, 24 insertions(+), 5 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index 1189727..ea9b51a 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -315,9 +315,9 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: shampoo_beta = 1 - (1 - shampoo_beta) / (1 - shampoo_beta**curr_iter_1_based) if state["step"] == 0: - preconditioner.init_step(grad, group["shampoo_beta"]) + preconditioner.init_step(grad, shampoo_beta) else: - preconditioner.step(grad, group["shampoo_beta"]) + preconditioner.step(grad, shampoo_beta) preconditioned_update = preconditioner.precondition(scalar_update) self._apply_weight_decay_inplace( diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 3378ad3..2390075 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -213,21 +213,24 @@ def __init__(self, state: dict, p_inv_root: float, eps: float) -> None: self.p_inv_root = p_inv_root self.eps = eps + # Store shampoo beta for verifing its value recieved in step. + self.shampoo_beta = None + @staticmethod def init_state(shape: tuple[int, ...], device: torch.device) -> dict[str, torch.Tensor]: return {} def rebind_state(self, state: dict) -> None: - pass + state["shampoo_beta"] = self.shampoo_beta def init_step(self, grad: torch.Tensor, shampoo_beta: float) -> None: - pass + self.shampoo_beta = shampoo_beta def update_kronecker_factors(self, grad: torch.Tensor, shampoo_beta: float) -> None: pass def step(self, grad: torch.Tensor, shampoo_beta: float) -> None: - pass + self.shampoo_beta = shampoo_beta def precondition(self, x: torch.Tensor) -> torch.Tensor: return x @@ -318,6 +321,22 @@ def test_3steps_smoke(self, m: int, n: int) -> None: self.assertEqual(state["step"], 3) self.assertCountEqual(state, {"step", "exp_avg", "L", "R"}) + def test_shampoo_beta_bias_corrected_over_5steps(self) -> None: + shampoo_beta = 0.75 + p = torch.nn.Parameter(torch.randn(4, 3, device=self.device)) + optimizer = _SgdShampoo([p], lr=1e-3, shampoo_beta=shampoo_beta) + + for curr_iter_1_based in range(1, 6): + p.grad = torch.randn_like(p) + optimizer.step() + + geometric_weight_sum = sum(shampoo_beta**i for i in range(curr_iter_1_based)) + self.assertEqual( + optimizer.state[p]["shampoo_beta"], + 1 - 1 / geometric_weight_sum, + msg=f"bias corrected shampoo_beta mismatch at step {curr_iter_1_based}", + ) + @parameterized.parameters(True, False) def test_init_group_skip_non_grad_params(self, skip_non_grad_params: bool) -> None: with_grad = torch.nn.Parameter(torch.randn(4, 3, device=self.device)) From 0a504e6b9d81f82c05fa0c812857b08b54e7695f Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 18 Aug 2026 13:07:29 -0700 Subject: [PATCH 13/16] add test for tikhonov Signed-off-by: Hao Wu --- tests/test_shampoo.py | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 2390075..832238c 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -15,7 +15,7 @@ from typing import override import torch -from _comparison import assert_equal +from _comparison import assert_close_to_identity, assert_equal from absl import flags, logging from absl.testing import absltest, parameterized @@ -152,6 +152,20 @@ def test_get_root_inverse_close_to_svd_reference(self, m: int, p_inv_root: int) rtol=1e-3, ) + @parameterized.parameters(2, 4) + def test_get_root_inverse_tikhonov_eps_effect(self, p_inv_root: int) -> None: + eps = 2.0**-4 + preconditioner = ShampooPreconditioner( + {"L": torch.eye(7, device=self.device), "R": torch.eye(7, device=self.device)}, + p_inv_root=p_inv_root, + eps=eps, + ) + + root_inverse = preconditioner._get_root_inverse(preconditioner.kronecker_factor_pair.L) + scale = 1 / (1 + eps ** (2 / p_inv_root)) + + assert_close_to_identity(root_inverse / scale) + @parameterized.parameters((6, 4), (4, 6), (5, 5)) def test_precondition_identity_factors_is_noop(self, m: int, n: int) -> None: preconditioner = ShampooPreconditioner( From 4b41c618575c8a26074d964ae11b9ddd0c682a45 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Wed, 19 Aug 2026 10:30:57 -0700 Subject: [PATCH 14/16] update tests Signed-off-by: Hao Wu --- tests/test_shampoo.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 832238c..6ee9fdf 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -309,9 +309,17 @@ def test_rejects_non_2d(self) -> None: with self.assertRaisesRegex(TypeError, "only supported for 2D"): optimizer.step() - def test_negative_lr_raises(self) -> None: - with self.assertRaisesRegex(ValueError, "Invalid learning rate"): - _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], lr=-1.0) + @parameterized.parameters( + {"kwargs": {"lr": -1.0}, "message": "Invalid learning rate"}, + {"kwargs": {"lr": 1e-3, "p_inv_root": -2}, "message": "p_inv_root must be positive integer"}, + ) + def test_invalid_arguments_raise(self, kwargs: dict, message: str) -> None: + with self.assertRaisesRegex(ValueError, message): + _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], **kwargs) + + @parameterized.parameters(2, 4.0) + def test_integral_p_inv_root_accepted(self, p_inv_root: float) -> None: + _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], lr=1e-3, p_inv_root=p_inv_root) class ShampooTest(parameterized.TestCase): From f7c6e0f8788e6bb743c6f9a1ee76c43f4e7e3a4d Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Wed, 19 Aug 2026 13:08:38 -0700 Subject: [PATCH 15/16] renamve to p_root_inv Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 22 ++++++------ tests/test_shampoo.py | 46 +++++++++++++------------- 2 files changed, 34 insertions(+), 34 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index ea9b51a..f1f396e 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -40,18 +40,18 @@ class ShampooPreconditioner: Args: state: Per-parameter optimizer state holding ``L`` and ``R``. - p_inv_root: Inverse root order; each factor is applied as ``A^(-1/p_inv_root)``. + p_root_inv: Inverse root order; each factor is applied as ``A^(-1/p_root_inv)``. eps: Floor on the eigenvalue magnitudes before inversion. """ def __init__( self, state: dict, - p_inv_root: float, + p_root_inv: float, eps: float, ) -> None: self.kronecker_factor_pair = precond_base.TensorPair(state["L"], state["R"]) - self.p_inv_root = p_inv_root + self.p_root_inv = p_root_inv self.eps = eps @staticmethod @@ -129,7 +129,7 @@ def step(self, grad: torch.Tensor, shampoo_beta: float) -> None: self.update_kronecker_factors(grad, shampoo_beta) def _get_root_inverse(self, kronecker_factor: torch.Tensor) -> torch.Tensor: - """Computes ``kronecker_factor^(-1/p_inv_root)`` from its eigendecomposition. + """Computes ``kronecker_factor^(-1/p_root_inv)`` from its eigendecomposition. Args: kronecker_factor: left or right kronecker factor @@ -143,7 +143,7 @@ def _get_root_inverse(self, kronecker_factor: torch.Tensor) -> torch.Tensor: eigvals = eigvals.clamp_min(0) # Tikhonov regularization - exp = 1.0 / self.p_inv_root + exp = 1.0 / self.p_root_inv inv_root_scale = eigvals**exp / (eigvals ** (2 * exp) + self.eps ** (2 * exp)) return (eigvecs * inv_root_scale) @ eigvecs.mT @@ -181,7 +181,7 @@ class ShampooBase(optim.Optimizer, opt_mixin.WeightDecayMixin): shampoo_beta: Kronecker factor EMA coefficient. eps: Numerical epsilon weight_decay: Decoupled weight decay coefficient. - p_inv_root: Inverse root order applied to each Kronecker factor. + p_root_inv: Inverse root order applied to each Kronecker factor. Attributes: PreconditionerCls: Preconditioner used for every parameter, and the source of the state layout @@ -199,16 +199,16 @@ def __init__( eps: float = 1e-8, weight_decay: float = 0.01, *, - p_inv_root: float = 4, + p_root_inv: float = 4, ) -> None: self.eps = eps self.weight_decay_method = "decoupled" - self.p_inv_root = p_inv_root + self.p_root_inv = p_root_inv if lr < 0.0: raise ValueError(f"Invalid learning rate: {lr}") - if p_inv_root <= 0 or round(p_inv_root) != p_inv_root: - raise ValueError(f"p_inv_root must be positive integer, got {p_inv_root}") + if p_root_inv <= 0 or round(p_root_inv) != p_root_inv: + raise ValueError(f"p_root_inv must be positive integer, got {p_root_inv}") defaults = { "lr": lr, @@ -305,7 +305,7 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: grad = p.grad.to(torch.float32) state = self.state[p] - preconditioner = self.PreconditionerCls(state, self.p_inv_root, self.eps) + preconditioner = self.PreconditionerCls(state, self.p_root_inv, self.eps) scalar_update = self._scalar_update(grad, state["exp_avg"], momentum=group["momentum"]) diff --git a/tests/test_shampoo.py b/tests/test_shampoo.py index 6ee9fdf..b36debf 100644 --- a/tests/test_shampoo.py +++ b/tests/test_shampoo.py @@ -37,9 +37,9 @@ def setUpModule() -> None: torch.cuda.manual_seed_all(FLAGS.seed) -def _root_inverse_reference(a: torch.Tensor, p_inv_root: float, eps: float) -> torch.Tensor: +def _root_inverse_reference(a: torch.Tensor, p_root_inv: float, eps: float) -> torch.Tensor: u, s, vh = torch.linalg.svd(a) - return (u * s.clamp_min(eps) ** (-1.0 / p_inv_root)) @ vh + return (u * s.clamp_min(eps) ** (-1.0 / p_root_inv)) @ vh def gen_signed_permutation(m: int): @@ -74,7 +74,7 @@ def test_init_state_rejects_non_2d(self) -> None: @parameterized.parameters((8, 16), (16, 8), (13, 15)) def test_rebind_state_binds_current_tensors(self, m: int, n: int) -> None: state = ShampooPreconditioner.init_state((m, n), self.device) - preconditioner = ShampooPreconditioner(state, p_inv_root=4, eps=1e-8) + preconditioner = ShampooPreconditioner(state, p_root_inv=4, eps=1e-8) preconditioner.step(torch.randn(m, n, device=self.device), 0.95) preconditioner.rebind_state(state) @@ -83,7 +83,7 @@ def test_rebind_state_binds_current_tensors(self, m: int, n: int) -> None: def test_rebind_state_missing_key_raises(self) -> None: state = ShampooPreconditioner.init_state((4, 4), self.device) - preconditioner = ShampooPreconditioner(state, p_inv_root=4, eps=1e-8) + preconditioner = ShampooPreconditioner(state, p_root_inv=4, eps=1e-8) del state["R"] with self.assertRaisesRegex(KeyError, "missing keys"): @@ -93,7 +93,7 @@ def test_rebind_state_missing_key_raises(self) -> None: def test_init_step_seeds_eps_identity(self, m: int, n: int) -> None: eps = 0.5 preconditioner = ShampooPreconditioner( - ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps + ShampooPreconditioner.init_state((m, n), self.device), p_root_inv=4, eps=eps ) grad = torch.zeros(m, n, device=self.device) @@ -106,7 +106,7 @@ def test_init_step_seeds_eps_identity(self, m: int, n: int) -> None: def test_update_kronecker_factors_matches_legacy(self, shape: tuple[int, int], shampoo_beta: float) -> None: m, n = shape preconditioner = ShampooPreconditioner( - ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=1e-8 + ShampooPreconditioner.init_state((m, n), self.device), p_root_inv=4, eps=1e-8 ) preconditioner.init_step(torch.randn(m, n, device=self.device), shampoo_beta) @@ -126,20 +126,20 @@ def test_step_equals_update_kronecker_factors(self) -> None: m, n, shampoo_beta = 6, 4, 0.9 grad = torch.randn(m, n, device=self.device) - stepped = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps) - updated = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=eps) + stepped = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_root_inv=4, eps=eps) + updated = ShampooPreconditioner(ShampooPreconditioner.init_state((m, n), self.device), p_root_inv=4, eps=eps) stepped.step(grad, shampoo_beta) updated.update_kronecker_factors(grad, shampoo_beta) assert_equal(stepped.kronecker_factor_pair.L, updated.kronecker_factor_pair.L) assert_equal(stepped.kronecker_factor_pair.R, updated.kronecker_factor_pair.R) - @parameterized.product(m=[4, 9, 16], p_inv_root=[2, 4]) - def test_get_root_inverse_close_to_svd_reference(self, m: int, p_inv_root: int) -> None: + @parameterized.product(m=[4, 9, 16], p_root_inv=[2, 4]) + def test_get_root_inverse_close_to_svd_reference(self, m: int, p_root_inv: int) -> None: x = 2 ** torch.randint(-3, 2, (m, m), device=self.device, dtype=torch.float) factor = x @ x.T + 0.125 * torch.eye(m, device=self.device) preconditioner = ShampooPreconditioner( - ShampooPreconditioner.init_state((m, m), self.device), p_inv_root=p_inv_root, eps=0 + ShampooPreconditioner.init_state((m, m), self.device), p_root_inv=p_root_inv, eps=0 ) with utils.fp32_matmul_precision("highest"): @@ -147,29 +147,29 @@ def test_get_root_inverse_close_to_svd_reference(self, m: int, p_inv_root: int) torch.testing.assert_close( root_inverse, - _root_inverse_reference(factor, p_inv_root, 0), + _root_inverse_reference(factor, p_root_inv, 0), atol=1e-3, rtol=1e-3, ) @parameterized.parameters(2, 4) - def test_get_root_inverse_tikhonov_eps_effect(self, p_inv_root: int) -> None: + def test_get_root_inverse_tikhonov_eps_effect(self, p_root_inv: int) -> None: eps = 2.0**-4 preconditioner = ShampooPreconditioner( {"L": torch.eye(7, device=self.device), "R": torch.eye(7, device=self.device)}, - p_inv_root=p_inv_root, + p_root_inv=p_root_inv, eps=eps, ) root_inverse = preconditioner._get_root_inverse(preconditioner.kronecker_factor_pair.L) - scale = 1 / (1 + eps ** (2 / p_inv_root)) + scale = 1 / (1 + eps ** (2 / p_root_inv)) assert_close_to_identity(root_inverse / scale) @parameterized.parameters((6, 4), (4, 6), (5, 5)) def test_precondition_identity_factors_is_noop(self, m: int, n: int) -> None: preconditioner = ShampooPreconditioner( - {"L": torch.eye(m, device=self.device), "R": torch.eye(n, device=self.device)}, p_inv_root=4, eps=0 + {"L": torch.eye(m, device=self.device), "R": torch.eye(n, device=self.device)}, p_root_inv=4, eps=0 ) x = torch.randn(m, n, device=self.device) @@ -193,7 +193,7 @@ def test_precondition_matches_inverse_of_known_spectrum(self, m: int) -> None: "R": A.clone(), } inv_root_kwargs = { - "p_inv_root": 2, + "p_root_inv": 2, "eps": 0, } preconditioner = ShampooPreconditioner(init_kronecker_factors, **inv_root_kwargs) @@ -211,7 +211,7 @@ def test_precondition_matches_inverse_of_known_spectrum(self, m: int) -> None: def test_precondition_4steps_smoke(self, m: int, n: int) -> None: shampoo_beta = 0.95 preconditioner = ShampooPreconditioner( - ShampooPreconditioner.init_state((m, n), self.device), p_inv_root=4, eps=1e-8 + ShampooPreconditioner.init_state((m, n), self.device), p_root_inv=4, eps=1e-8 ) preconditioner.init_step(torch.randn(m, n, device=self.device), shampoo_beta) for _ in range(4): @@ -223,8 +223,8 @@ def test_precondition_4steps_smoke(self, m: int, n: int) -> None: class _BypassPreconditioner: - def __init__(self, state: dict, p_inv_root: float, eps: float) -> None: - self.p_inv_root = p_inv_root + def __init__(self, state: dict, p_root_inv: float, eps: float) -> None: + self.p_root_inv = p_root_inv self.eps = eps # Store shampoo beta for verifing its value recieved in step. @@ -311,15 +311,15 @@ def test_rejects_non_2d(self) -> None: @parameterized.parameters( {"kwargs": {"lr": -1.0}, "message": "Invalid learning rate"}, - {"kwargs": {"lr": 1e-3, "p_inv_root": -2}, "message": "p_inv_root must be positive integer"}, + {"kwargs": {"lr": 1e-3, "p_root_inv": -2}, "message": "p_root_inv must be positive integer"}, ) def test_invalid_arguments_raise(self, kwargs: dict, message: str) -> None: with self.assertRaisesRegex(ValueError, message): _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], **kwargs) @parameterized.parameters(2, 4.0) - def test_integral_p_inv_root_accepted(self, p_inv_root: float) -> None: - _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], lr=1e-3, p_inv_root=p_inv_root) + def test_integral_p_root_inv_accepted(self, p_root_inv: float) -> None: + _SgdShampoo([torch.nn.Parameter(torch.randn(4, 4, device=self.device))], lr=1e-3, p_root_inv=p_root_inv) class ShampooTest(parameterized.TestCase): From 15fa3a155b7ec46af0ec2713f595e3a2708c3b8c Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Wed, 19 Aug 2026 13:20:51 -0700 Subject: [PATCH 16/16] rename invere root to root inverse Signed-off-by: Hao Wu --- emerging_optimizers/shampoo/shampoo.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/emerging_optimizers/shampoo/shampoo.py b/emerging_optimizers/shampoo/shampoo.py index f1f396e..c4fa290 100644 --- a/emerging_optimizers/shampoo/shampoo.py +++ b/emerging_optimizers/shampoo/shampoo.py @@ -148,7 +148,7 @@ def _get_root_inverse(self, kronecker_factor: torch.Tensor) -> torch.Tensor: return (eigvecs * inv_root_scale) @ eigvecs.mT def precondition(self, x: torch.Tensor) -> torch.Tensor: - """Applies both inverse roots to a matrix in the parameter basis. + """Applies both root inverse to a matrix in the parameter basis. Args: x: Matrix in the parameter basis. @@ -156,12 +156,10 @@ def precondition(self, x: torch.Tensor) -> torch.Tensor: Returns: The preconditioned matrix, in the parameter basis. """ - inverse_root_pair = precond_base.TensorPair( - self._get_root_inverse(self.kronecker_factor_pair.L), - self._get_root_inverse(self.kronecker_factor_pair.R), - ) + root_inv_L = self._get_root_inverse(self.kronecker_factor_pair.L) + root_inv_R = self._get_root_inverse(self.kronecker_factor_pair.R) - return inverse_root_pair.L @ x @ inverse_root_pair.R + return root_inv_L @ x @ root_inv_R class ShampooBase(optim.Optimizer, opt_mixin.WeightDecayMixin):