import torch

@torch.compile
def zeropower_via_newtonschulz5(G: torch.Tensor, steps: int = 5) -> torch.Tensor:
    """
    Muon 5-step Newton-Schulz polynomial matrix orthogonalization kernel.
    Computes nearest semi-orthogonal matrix Ortho(G) = UV^T via native bfloat16 matmuls.
    
    Polynomial coefficients:
      phi(x) = 3.4445*x - 4.7750*x^3 + 2.0315*x^5
    """
    assert G.ndim == 2, "Muon requires 2D weight matrix gradients"
    a, b, c = (3.4445, -4.7750, 2.0315)
    X = G.bfloat16()
    X /= (X.norm() + 1e-7)
    
    if G.size(0) > G.size(1):
        X = X.T

    for _ in range(steps):
        A = X @ X.T
        B = b * A + c * A @ A
        X = a * X + B @ X

    if G.size(0) > G.size(1):
        X = X.T
        
    return X.to(G.dtype)


class Muon(torch.optim.Optimizer):
    """
    Muon Matrix Optimizer for 2D weight matrices in deep neural networks.
    """
    def __init__(self, params, lr: float = 0.02, momentum: float = 0.95, nesterov: bool = True, ns_steps: int = 5):
        defaults = dict(lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps)
        super().__init__(params, defaults)

    @torch.no_grad()
    def step(self):
        for group in self.param_groups:
            lr = group['lr']
            momentum = group['momentum']
            for p in group['params']:
                if p.grad is None:
                    continue
                g = p.grad
                state = self.state[p]
                if 'momentum_buffer' not in state:
                    state['momentum_buffer'] = torch.zeros_like(g)
                buf = state['momentum_buffer']
                buf.mul_(momentum).add_(g)
                if group['nesterov']:
                    g = g.add(buf, alpha=momentum)
                else:
                    g = buf
                update = zeropower_via_newtonschulz5(g, steps=group['ns_steps'])
                p.data.add_(update, alpha=-lr)
