Weight decay controls delayed learning and generalization in neural nets

A Spectral Theory of Grokking: Weight Decay induces Feature Learning

Machine LearningArtificial Intelligence

Summary

Sometimes neural networks learn to fit training data quickly but only get better at generalizing to new data much later, a phenomenon called grokking. The authors explain this delay by showing how weight decay, a common training technique, causes the network to slowly update feature representations after initially memorizing the training set. They provide a mathematical model predicting how learning rate and weight decay combine to influence the timing and success of this transition from simple memorization to richer understanding. The theory is supported by experiments on modular addition tasks using multilayer perceptrons and transformers. This helps understand why and when neural networks start to truly generalize during training.

What this means in practice

  • For machine learning engineers: Adjust weight decay and learning rate together to control when neural networks begin effective feature learning and improve generalization.
  • For ai system trainers: Diagnose and avoid training regimes where excessive weight decay prevents fitting, saving computational resources on invalid model runs.

Tested on simulated data.

Authors

Lenz Pracher, Pascal de Jong, Oskar Lieshaus, Alan Jeffares, Steffen Rulands

Abstract

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.