Summary (Overview)

  • Key contribution: This paper demonstrates that Muon optimizer "sublates" the classical edge-of-stability (EoS) phenomenon in LLM pretraining—it preserves a stochastic loss-neutral edge but breaks the classical coupling between loss balance and update reversal, resulting in a "split edge of stability."
  • Theoretical result: The authors derive a coherence-corrected conditional loss-neutral boundary sk,bM=2ρk,b/ηks_{k,b}^{M} = 2\rho_{k,b}/\eta_k for stochastic no-momentum Muon, generalizing the full-batch GD boundary 2/η2/\eta.
  • Key finding: Loss balance (T1) and temporal alignment (T2) are distinct diagnostics that respond differently to learning rate and batch size; negative temporal alignment does not determine which side of the loss-neutral boundary an update lies on.
  • Empirical validation: Across controlled experiments (MLP, CNN, Transformer) and LLM pretraining (22M, 130M, 1B parameters), larger batch sizes drive temporal alignment toward coherent reversal (cpol=−1c^{pol} = -1), while small batches remain near orthogonality.
  • Practical implication: The two signals (T1 and T2) should be monitored jointly when designing learning-rate and batch-size schedules; a larger peak learning rate can improve validation loss while cpolc^{pol} remains controlled.

Introduction and Theoretical Foundation

The paper addresses a fundamental gap in understanding Muon's large-step dynamics during LLM pretraining. Classical gradient descent (GD) exhibits the edge of stability (EoS) phenomenon where curvature approaches a learning-rate-dependent boundary, loss becomes non-monotone over short horizons, yet training progresses over longer ones. In a scalar quadratic mode of curvature λ\lambda, the GD multiplier is 1−ηλ1 - \eta\lambda, and at ηλ=2\eta\lambda = 2, three signatures coincide: one-step loss neutrality, equal-magnitude reversal, and marginal linear stability.

Key theoretical distinction: Muon replaces each matrix gradient with an approximately semi-orthogonal polar direction, which fundamentally changes the geometry of the update. The paper shows that Muon removes the mechanism forcing these EoS signatures to coincide because:

  • Its loss change depends on curvature along the current polar direction
  • Its temporal geometry depends on how that direction changes between updates
  • Stochasticity introduces a further distinction: an unbiased minibatch gradient need not induce an unbiased polar direction

Muon formulation: For a continuously differentiable objective L:Rm×n→RL: \mathbb{R}^{m\times n} \to \mathbb{R} with batch loss LBL_B and fresh batch Bk∼DbB_k \sim \mathcal{D}_b of size bb:

Gk=∇L(Wk),G^k=∇LBk(Wk),Pk=Pol(G^k),Wk+1=Wk−ηPk\boldsymbol{G}_k = \nabla L(\boldsymbol{W}_k), \quad \widehat{\boldsymbol{G}}_k = \nabla L_{B_k}(\boldsymbol{W}_k), \quad \boldsymbol{P}_k = \text{Pol}(\widehat{\boldsymbol{G}}_k), \quad \boldsymbol{W}_{k+1} = \boldsymbol{W}_k - \eta \boldsymbol{P}_k

where Pol(G)=UrVr⊤\text{Pol}(G) = U_r V_r^\top from the compact SVD G=UrΣrVr⊤G = U_r \Sigma_r V_r^\top. The operator–nuclear duality gives:

max⁡∥Z∥op≤1⟨G,Z⟩F=⟨G,Pol(G)⟩F=∥G∥∗\max_{\|\boldsymbol{Z}\|_{op} \leq 1} \langle \boldsymbol{G}, \boldsymbol{Z}\rangle_F = \langle \boldsymbol{G}, \text{Pol}(\boldsymbol{G})\rangle_F = \|\boldsymbol{G}\|_*

Methodology

Key Diagnostics

The paper defines two primary diagnostics:

T1 (Conditional loss neutrality):

sk,bM=2ρk,bηks_{k,b}^{M} = \frac{2\rho_{k,b}}{\eta_k}

where ρb\rho_b is the batch-to-population coherence and sbMs_b^M is the effective curvature.

T2 (Temporal orthogonality):

ck,bpol=⟨Pk,Pk+1⟩F∥Pk∥F∥Pk+1∥Fc_{k,b}^{pol} = \frac{\langle \boldsymbol{P}_k, \boldsymbol{P}_{k+1}\rangle_F}{\|\boldsymbol{P}_k\|_F \|\boldsymbol{P}_{k+1}\|_F}

Theoretical Framework

Definition 2.1 (Coherence and effective curvature):

ρb(Wk):=Eb⟨Gk,Pk⟩F∥Gk∥∗∈[−1,1]\rho_b(\boldsymbol{W}_k) := \frac{\mathbb{E}_b \langle \boldsymbol{G}_k, \boldsymbol{P}_k\rangle_F}{\|\boldsymbol{G}_k\|_*} \in [-1, 1] sbM(Wk;η):=2Eb[L(Wk−ηPk)−L(Wk)+η⟨Gk,Pk⟩F]η2∥Gk∥∗→η→0Eb⟨Pk,Hk[Pk]⟩F∥Gk∥∗s_b^M(\boldsymbol{W}_k; \eta) := \frac{2\mathbb{E}_b[L(\boldsymbol{W}_k - \eta\boldsymbol{P}_k) - L(\boldsymbol{W}_k) + \eta\langle\boldsymbol{G}_k, \boldsymbol{P}_k\rangle_F]}{\eta^2 \|\boldsymbol{G}_k\|_*} \xrightarrow{\eta \to 0} \frac{\mathbb{E}_b\langle\boldsymbol{P}_k, \boldsymbol{H}_k[\boldsymbol{P}_k]\rangle_F}{\|\boldsymbol{G}_k\|_*}

Theorem 2.1 (Conditional loss identity):

Eb[L(Wk+1)−L(Wk)∣Wk]=η2∥Gk∥∗2(sbM(Wk)−2ρb(Wk)η)\mathbb{E}_b[L(\boldsymbol{W}_{k+1}) - L(\boldsymbol{W}_k) \mid \boldsymbol{W}_k] = \frac{\eta^2 \|\boldsymbol{G}_k\|_*}{2}\left(s_b^M(\boldsymbol{W}_k) - \frac{2\rho_b(\boldsymbol{W}_k)}{\eta}\right)

The conditional expected loss is nonincreasing if and only if sbM≤2ρb/ηs_b^M \leq 2\rho_b/\eta. At full batch, ρfull=1\rho_{full} = 1, recovering the classical 2/η2/\eta boundary.

Experimental Setup

  • Controlled experiments: MLP, CNN, and Transformer with SVD-polar updates at fixed learning rate η=0.007\eta = 0.007 with varying batch sizes (64–512 and full batch)
  • LLM experiments:
    • 130M Llama-like LLM on FineWeb with batch size 32, sequence length 4,096, processing 5.23B tokens
    • 22M model for perturbation experiments
    • 1B Llama-like model with batch size 512, using 20B tokens
  • Optimizer: No-momentum NS-5 Muon approximation on Muon parameter blocks, with AdamW on auxiliary parameters
  • Learning rate schedule: Warmup-stable-decay (WSD) with peak learning rates 0.02–0.12

Empirical Validation / Results

Controlled Network Experiments

Loss boundary tracking: Across MLP, CNN, and Transformer experiments, the conditional curvature s^bM\hat{s}_b^M approaches and oscillates around the coherence-corrected boundary 2ρ^b/η2\hat{\rho}_b/\eta, marking near-zero conditional loss increments and sign changes. Larger batch sizes encourage larger ρb\rho_b and induce less fluctuation.

Temporal alignment vs. batch size: Increasing batch size makes temporal alignment progressively more negative:

  • Full-batch trajectories move toward coherent reversal at cpol=−1c^{pol} = -1
  • Small-batch trajectories remain much closer to orthogonality

Order of T1/T2 events: No consistent ordering exists—at b=512b = 512, T2 precedes T1 in MLP and CNN, while at b=4096b = 4096, T1 precedes T2 (Figure 3).

Toy Model: Quadratic Linear Model

For the matrix quadratic LX,E(W)=12∥WX−Y∥F2L_{\boldsymbol{X},\boldsymbol{E}}(\boldsymbol{W}) = \frac{1}{2}\|\boldsymbol{W}\boldsymbol{X} - \boldsymbol{Y}\|_F^2 with Y=W⋆X+E\boldsymbol{Y} = \boldsymbol{W}_\star\boldsymbol{X} + \boldsymbol{E}, the thresholds are:

skM=∥X∥F2tr(Sk),ckpol=1−2#n−(Sk−ηXX⊤)ds_k^M = \frac{\|\boldsymbol{X}\|_F^2}{\text{tr}(\boldsymbol{S}_k)}, \qquad c_k^{pol} = 1 - \frac{2\#_{n_-}(\boldsymbol{S}_k - \eta\boldsymbol{X}\boldsymbol{X}^\top)}{d}

where #n−(A)\#_{n_-}(A) counts negative eigenvalues. The paper shows that curvature depends on a trace while alignment depends on a sign count—demonstrating that curvature alone cannot determine alignment, even at the same learning rate and full coherence.

130M LLM Pretraining Results

  • Loss balance coexists with continued learning: Conditional curvature repeatedly lies near 2ρ^k,32/ηk2\hat{\rho}_{k,32}/\eta_k, indicating near-zero conditional mean reference-loss increments, yet validation loss improves over longer horizons
  • Global directions stay near orthogonality: Full-horizon cosine medians range from −0.048 to −0.032, far from exact reversal at −1
  • Negative alignment varies across layers: Middle-layer key and value matrices have more negative cosine medians than many query matrices (Figure 6)

Learning-Rate Perturbation Results (22M and 130M)

  • Alignment tends to return toward its pre-perturbation value: Increasing learning rate initially makes cpolc^{pol} more negative; decreasing it makes cpolc^{pol} less negative; both tend to recover
  • Loss balance and alignment respond differently: Halving the learning rate puts the conditional step on the loss-decrease side of T1; doubling it puts the step on the loss-increase side, yet temporal cosines remain negative in all branches
  • Validation loss improves under larger peak learning rate while cpolc^{pol} remains controlled, suggesting the possibility of increasing peak LR in WSD guided by cpolc^{pol}

1B LLM Pretraining Results

  • Stronger partial cancellation: The 1B run (batch size 512) shows more negative global alignment cpolc^{pol} than 130M (batch size 32), consistent with the batch-size trend
  • Still far from coherent reversal: cpolc^{pol} remains far from −1 despite stronger cancellation
  • Layer heterogeneity: Layer-level cosines vary along the trajectory; a single global value does not describe every layer

Theoretical and Practical Implications

Theoretical Implications

  1. Split edge of stability: The paper establishes that Muon's EoS is fundamentally different from GD's—the loss-neutral edge survives but is not accompanied by a universal temporal-direction signature. This breaks the classical coupling where loss neutrality, equal-magnitude reversal, and marginal stability coincide.

  2. Coherence-corrected boundary: The theoretical derivation shows that minibatch sampling moves the loss-neutral boundary through the batch-to-population coherence ρk,b\rho_{k,b}, generalizing the full-batch 2/η2/\eta result. This provides a rigorous foundation for understanding stochastic Muon dynamics.

  3. Geometry-dependent temporal boundary: For full-batch exact-polar Muon, the T2 condition is s‾kM=ϑk/η\overline{s}_k^M = \vartheta_k/\eta where ϑk\vartheta_k depends on the singular-value geometry of consecutive gradients—not a universal 1/η1/\eta threshold as in GD.

Practical Implications

  1. Joint monitoring of T1 and T2: The two signals should be monitored together during training. During WSD schedules, weaker negative alignment indicates reduced cancellation, but increasing the learning rate is justified only if the T1 diagnostic leaves sufficient room before loss neutrality.

  2. Data-driven critical batch size: Batch size affects both batch-to-population coherence and temporal alignment, suggesting a critical batch size based on their joint response rather than gradient noise alone.

  3. Learning rate scheduling guidance: The WSD perturbation experiments show that a larger peak learning rate can improve validation loss while cpolc^{pol} remains controlled, motivating schedules that monitor both quantities rather than treating negative alignment alone as instability.

Conclusion

This paper demonstrates that Muon does not inherit GD's EoS as a single edge: conditional loss balance and temporal direction become separate signals. Key findings include:

  • Controlled experiments: Increasing batch size drives cpolc^{pol} toward coherent reversal, while small batches remain closer to orthogonality
  • LLM scaling: The 130M runs track the loss-neutral boundary with weak negative alignment; the 1B run exhibits stronger partial cancellation but remains far from reversal
  • Stochastic LLM pretraining can reach the loss edge without reproducing the full-batch directional regime

Future directions: The paper notes that momentum introduces temporal filtering and an additional optimizer state, so the appropriate diagnostics must be reconsidered for coupled parameter–momentum dynamics. Extending the split-edge picture to momentum Muon is identified as an important direction for future work. Converting the observed distinct responses of T1 and T2 into optimized adaptive schedules is also left for future research.

Related papers