Post-Grokking Collapse at the Representation-Readout Interface in Muon-Trained Transformers
AuthorsAli Janati, Kaoutar El Maghraoui, Andrei Kanavalau, Anass Belfatmi
Resources
The study shows that Muon can make transformers grok quickly but later destabilize their representation-to-output interface, causing learned solutions to collapse.
Key results
Average steps to sustained 95% generalization across successful Muon configurations.
Average steps to sustained 95% generalization across successful AdamW configurations.
Post-grokking steps completed with zero evaluations below 95% after freezing embeddings and readout.
Accuracy after rescaling the task-aligned Fourier family without retraining.
What the paper found
This study examines post-grokking collapse in a decoder-only transformer trained on modular addition, (a+b) mod 113, with Muon updating hidden matrices and AdamW updating embeddings and the output head. Muon reaches sustained 95% generalization faster than AdamW—13011 steps on average versus 30486—but all nine Muon configurations subsequently lose generalization. The proposed mechanism is hidden–readout misalignment: once the training loss is near zero, the residual stream and unembedding are identifiable only jointly up to an invertible basis change, while Muon and AdamW respond differently to tiny gradients. Applied-step elasticity is −0.03 for the Muon group versus +1.5 for AdamW groups, producing parameter displacement separation at 8.0 times the rate. The Fourier circuit itself can remain intact: during one collapse, test accuracy falls to 19.04%, while dominant-frequency support retains a Jaccard index of 1.0000 and power retains cosine similarity 0.9899. The paper distinguishes circuit failure from circuit masking; in masking, the task-aligned Fourier family still achieves 100% alone while the full model reaches 45.85%, and rescaling that family restores accuracy to 99.9% without retraining. Freezing token and positional embeddings plus the readout after circuit formation prevents collapse across 451400 post-grokking steps, with zero evaluations below 95%. Removing Muon’s normalization and orthogonalization instead reduces spectral dispersion from 326 effective conjugate pairs to 4.11 and ends in non-finite loss, so it is not a stable replacement.
Original abstract
Under the standard split, Muon gets hidden matrices and AdamW embeddings/output head. Muon groks modular addition faster, but its solutions do not hold. All nine configurations on $(a+b) \bmod 113$ grok and later lose generalization. Across five seeds the selected AdamW reference falls below threshold on four, reaching 27.59%. Instability persists across two moduli, two widths, two training fractions, subtraction, and depth. The failure arises at the representation-readout interface, identified only jointly up to an invertible map unselected by the loss. After solving the training set, the gradient falls to order $10^{-6}$ and the optimizers respond differently: step-size elasticity is -0.03 for Muon versus +1.5 for AdamW, and the Muon group moves 8.0 times faster per parameter. From bit-identical states, freezing either group prevents failure. Freezing embeddings/readout removes it in five runs over 451,400 post-grokking steps and five paired seeds: unfrozen arms record 137-321 sub-threshold evaluations, frozen arms none. Removing Muon's normalization and orthogonalization is no substitute: it collapses representation from 326 effective conjugate pairs to 4, shows no recurrent collapse, and fails terminally. Fourier filtering separates circuit failure from masking. Across 43 checkpoints over five seeds and three regimes, the task-aligned family reaches exactly 100% alone. In circuit failure it no longer solves the task; in masking it remains perfect while the full model reaches 45.85%, giving a positive margin on every example, including errors, but being outvoted by a near-equal adversarial remainder. Rescaling it restores 99.9%; grokking is the same condition resolving upward. The task selects the family, swapping $(k,k)$ for $(k,-k)$ under subtraction. Across an abrupt collapse, standard Fourier support is unchanged and the power-distribution cosine remains 0.9899.
Read the original paperMore in Optimization
Browse all 36 papers →An $Ω(κ_y^8ε^{-6})$ Lower Bound for Stochastic NC-SC Bilevel Optimization with First-order Oracles
Zhihao Gu, Qilong Wu, Junchi Yang
This work proves that stochastic bilevel optimization fundamentally requires up to epsilon^{-6} oracle queries, showing existing methods are asymptotically optimal.
Hyper Algorithm Design Agent: Evolving Learnable Optimizer from Zero
Zipei Yu, Yue-Jiao Gong, Zeyuan Ma, Yuncheng Jiang, Zhiguang Cao
A pair of self-improving coding agents evolves new learnable optimization algorithms from a simple template, reducing the need for handcrafted optimizer design.
Tight Regret Bound for Online Inverse Linear Optimization via Multiscale Matrix Weights
Shinsaku Sakaue
A new multiscale matrix-weights algorithm learns hidden linear preferences online with provably optimal dimension-dependent regret.