A Spectral Theory of Grokking: Weight Decay induces Feature Learning
Paper derives a spectral theory showing weight decay drives NTK feature learning that controls grokking timescale and phase structure.
The paper provides a quantitative theory of grokking in homogeneous networks trained with squared loss and L2 weight decay, showing residuals after memorization feed back into NTK dynamics. It predicts the grokking timescale is controlled by the product of learning rate and weight decay, with feature learning slowing logarithmically near a critical decay threshold. Predictions were validated on modular addition using an 84x90 grid of trained MLPs and a 42x45 grid with a one-block Transformer, both recovering the predicted phase geometry and inverse-product scaling. Task-aligned Fourier structure continued emerging in the NTK after training accuracy saturated.
- Theory links weight decay-driven NTK evolution to delayed generalization in grokking
- Grokking timescale scales inversely with product of learning rate and weight decay
- Validated on 84x90 MLP grid and 42x45 Transformer grid in modular addition
- Critical decay threshold above which task-aligned NTK structure cannot generalize
Full article276 words · extracted from arxiv.org · click to collapse
In grokking an early fit to the training data separates from a much later improvement in generalization. During this delay, training can move from a fixed neural tangent kernel (NTK) regime to one in which task-relevant kernel eigendirections continue to evolve. We provide a quantitative theory for how this transition from lazy to rich learning can produce delayed generalization. For homogeneous networks trained with squared loss and $L_2$ weight decay, we show that a finite residual remains after memorization, with larger residual fractions in target components associated with smaller NTK eigenvalues. These residuals feed back into the dynamics of the NTK itself, and projecting the resulting dynamics onto task-relevant spectral directions yields a reduced system in which residual-driven kernel growth competes with weight decay. This system predicts that the grokking timescale is controlled by the product of learning rate and weight decay, that feature learning slows logarithmically near a critical decay above which task-aligned NTK structure can no longer support generalization, and that stronger decay can prevent fitting altogether. We test these predictions in modular addition. In a homogeneous MLP, task-aligned Fourier structure continues to emerge in the NTK after training accuracy has saturated, and an 84$\times$90-grid of trained networks across varying learning rate and weight decay recovers the predicted phase geometry and inverse-product scaling of the generalization time with learning rate and weight decay. A one-block Transformer shows similar macroscopic phase structure in a 42$\times$45-grid, as well as the same transition-time scaling despite violating exact homogeneity. Together, these results provide a mechanistic derivation connecting post-fit feature learning to both the onset of generalization and its phase structure in the learning rate and weight decay plane.
Text extracted automatically; images, tables and formatting may be missing. Original: https://arxiv.org/abs/2609.26679