From c62fddf5a4c6f5faecf7d6af0236f778243ad78b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 17 Mar 2026 15:52:21 -0700 Subject: [PATCH 1/5] add _init_group to adaptive muon Signed-off-by: Hao Wu --- .../adaptive_muon.py | 62 ++++++++++--------- 1 file changed, 34 insertions(+), 28 deletions(-) diff --git a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py index ae26dd98..cd9071c3 100644 --- a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py +++ b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py @@ -98,37 +98,45 @@ def __init__( group.setdefault("beta2", beta2) group.setdefault("eps", eps) - def _initialize_moment2( + @torch.no_grad() # type: ignore[misc] + @override + def _init_group( self, - state: dict[str, torch.Tensor], - grad: torch.Tensor, + group: dict, + skip_non_grad_params: bool = True, ) -> None: - """Initialize the second moment buffer if it doesn't exist. + """Performs lazy state initialization for parameters. - The shape of the buffer depends on the moment2_method: - - "adamuon": Full elementwise buffer with same shape as grad + Extends the base class to also initialize the second moment buffer. + The shape of the moment2 buffer depends on the moment2_method: + - "adamuon": Full elementwise buffer with same shape as parameter - "normuon": Reduced shape buffer (averaged along -1 if shape[-2] >= shape[-1], else -2) Args: - state: The optimizer state dict for a parameter. - grad: The gradient tensor (used for shape/dtype). + group: Parameter group dictionary. + skip_non_grad_params: If True, skip parameters without gradients. """ - if "moment2_buffer" not in state: - if self.moment2_method == "adamuon": - # Full elementwise second moment - moment2 = torch.zeros_like(grad) - elif self.moment2_method == "normuon": - # Row/column-wise second moment - reduced along one dimension - # Determine which dimension to reduce based on parameter shape - avg_dim = -1 if grad.shape[-2] >= grad.shape[-1] else -2 - # Specify the shape with reduced dimension - moment2_shape = list(grad.shape) - moment2_shape[avg_dim] = 1 - moment2 = torch.zeros(moment2_shape, dtype=grad.dtype, device=grad.device) - else: - raise TypeError(f"Invalid second moment method: {self.moment2_method}") - - state["moment2_buffer"] = moment2 + for p in group["params"]: + if skip_non_grad_params and p.grad is None: + continue + state = self.state[p] + + if len(state) == 0: + state["momentum_buffer"] = torch.zeros_like(p.data) + + if self.moment2_method == "adamuon": + # Full elementwise second moment + state["moment2_buffer"] = torch.zeros_like(p.data) + elif self.moment2_method == "normuon": + # Row/column-wise second moment - reduced along one dimension + # Determine which dimension to reduce based on parameter shape + avg_dim = -1 if p.data.shape[-2] >= p.data.shape[-1] else -2 + # Specify the shape with reduced dimension + moment2_shape = list(p.data.shape) + moment2_shape[avg_dim] = 1 + state["moment2_buffer"] = torch.zeros(moment2_shape, dtype=p.data.dtype, device=p.data.device) + else: + raise TypeError(f"Invalid second moment method: {self.moment2_method}") def _apply_moment2_normalization( self, @@ -204,6 +212,8 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: loss = closure() for group in self.param_groups: + self._init_group(group) + for p in group["params"]: if p.dim() != 2: raise ValueError(f"{self.__class__.__name__} only supports 2D parameters") @@ -212,10 +222,6 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None: continue state = self.state[p] - if "momentum_buffer" not in state: - state["momentum_buffer"] = torch.zeros_like(grad) - self._initialize_moment2(state, grad) - exp_avg = state["momentum_buffer"] self._apply_weight_decay_inplace( From 33782dafbebc4df70af89d852cca23ad3dc7ba60 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 17 Mar 2026 15:53:12 -0700 Subject: [PATCH 2/5] add experimental warning Signed-off-by: Hao Wu --- emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py index cd9071c3..93338ace 100644 --- a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py +++ b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py @@ -41,6 +41,9 @@ class AdaptiveMuon(muon.Muon): descent for deep learning.* In Advances in neural information processing systems 28 (2015). The step() method is overridden to include second moment normalization logic. + Warning: + This optimizer is experimental and may change in future versions. + Args: params: Iterable of parameters to optimize or dicts defining parameter groups. lr: Learning rate. From 69627bad9bd00b72c66067c9b4fa44f0258d79b9 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 17 Mar 2026 17:14:42 -0700 Subject: [PATCH 3/5] Update emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Signed-off-by: Hao Wu --- .../orthogonalized_optimizers/adaptive_muon.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py index 93338ace..23cd1d2d 100644 --- a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py +++ b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py @@ -133,6 +133,11 @@ def _init_group( elif self.moment2_method == "normuon": # Row/column-wise second moment - reduced along one dimension # Determine which dimension to reduce based on parameter shape + if p.data.ndim < 2: + raise ValueError( + f"{self.__class__.__name__} only supports 2D parameters, " + f"got shape {tuple(p.data.shape)}" + ) avg_dim = -1 if p.data.shape[-2] >= p.data.shape[-1] else -2 # Specify the shape with reduced dimension moment2_shape = list(p.data.shape) From 4dee92b0666f1960f9d3af38f9fc619904378c5c Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 17 Mar 2026 17:23:13 -0700 Subject: [PATCH 4/5] fix AI format error Signed-off-by: Hao Wu --- .../orthogonalized_optimizers/adaptive_muon.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py index 23cd1d2d..1974c1e7 100644 --- a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py +++ b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py @@ -133,10 +133,9 @@ def _init_group( elif self.moment2_method == "normuon": # Row/column-wise second moment - reduced along one dimension # Determine which dimension to reduce based on parameter shape - if p.data.ndim < 2: + if p.data.ndim != 2: raise ValueError( - f"{self.__class__.__name__} only supports 2D parameters, " - f"got shape {tuple(p.data.shape)}" + f"{self.__class__.__name__} only supports 2D parameters, got shape {tuple(p.data.shape)}" ) avg_dim = -1 if p.data.shape[-2] >= p.data.shape[-1] else -2 # Specify the shape with reduced dimension From e06ba8570866a46c391c032d6143420edcd4707b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Tue, 17 Mar 2026 17:23:31 -0700 Subject: [PATCH 5/5] fix AI format error Signed-off-by: Hao Wu --- emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py index 1974c1e7..05ac8c2a 100644 --- a/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py +++ b/emerging_optimizers/orthogonalized_optimizers/adaptive_muon.py @@ -135,7 +135,7 @@ def _init_group( # Determine which dimension to reduce based on parameter shape if p.data.ndim != 2: raise ValueError( - f"{self.__class__.__name__} only supports 2D parameters, got shape {tuple(p.data.shape)}" + f"{self.__class__.__name__} only supports 2D parameters, got shape {tuple(p.data.shape)}" ) avg_dim = -1 if p.data.shape[-2] >= p.data.shape[-1] else -2 # Specify the shape with reduced dimension