""" Muon optimizer -- MomentUm Orthogonalized by Newton-Schulz. https://kellerjordan.github.io/posts/muon/ Included locally so this repository is self-contained for anyone who wants to fine-tune or continue pretraining from the exported checkpoint. """ import torch from torch import Tensor @torch.compile def zeropower_via_newtonschulz5(G: Tensor, steps: int) -> Tensor: assert G.ndim >= 2 a, b, c = (3.4445, -4.7750, 2.0315) X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7) for _ in range(steps): A = X @ X.mT B = b * A + c * A @ A X = a * X + B @ X if G.size(-2) > G.size(-1): X = X.mT return X class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, momentum=0.95, nesterov=True, ns_steps=5): defaults = dict(lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps) params = [*params] param_groups = [] for size in {p.numel() for p in params}: param_groups.append(dict(params=[p for p in params if p.numel() == size])) super().__init__(param_groups, defaults) @torch.no_grad() def step(self): for group in self.param_groups: for p in group["params"]: g = p.grad assert g is not None state = self.state[p] if "momentum_buffer" not in state: state["momentum_buffer"] = torch.zeros_like(g) buf = state["momentum_buffer"] buf.lerp_(g, 1 - group["momentum"]) g = g.lerp_(buf, group["momentum"]) if group["nesterov"] else buf g = zeropower_via_newtonschulz5(g, steps=group["ns_steps"]) p.add_(g, alpha=-group["lr"] * max(1, p.size(-2) / p.size(-1)) ** 0.5)