High-dimensional Asymptotics of Feature Learning: How One Gradient Step Improves the Representation
Jimmy BaMurat A. ErdogduTaiji SuzukiZhichao WangDenny WuGreg Yang
Proves that taking a single gradient descent step on the first layer of a two-layer neural network with a sufficiently large learning rate enables the resulting kernel to outperform fixed random features and surpass linear estimators in high dimensions.
Modern machine learning increasingly relies on deep neural networks because they automatically extract and adapt useful internal representations from data, a capability known as feature learning. In contrast, traditional kernel methods and fixed random feature models rely on static, unlearned transformations that often perform no better than simple linear predictors in high-dimensional tasks. While empirical observations show that critical representation learning occurs in the very earliest training iterations, theoretical frameworks have struggled to precisely capture when, how, and by how much gradient-based training breaks free from the limitations of fixed kernels.
The article establishes a precise mathematical framework to evaluate how taking a single gradient descent step on the first layer of a two-layer neural network improves representation quality and prediction risk. Specifically, it analyzes whether one feature learning step enables the model to outperform fixed kernel benchmarks in high-dimensional settings where the dataset size, input dimensionality, and hidden neuron count scale proportionally.
To conduct this evaluation, the authors combine random matrix theory and operator-valued free probability with high-dimensional statistical simulations. They study a student-teacher setup where the teacher is a single-index model with Gaussian inputs. The first-layer network weights undergo one step of gradient descent on a training set, and the quality of the updated features is then evaluated by computing the prediction risk of kernel ridge regression trained on a fresh sample of independent data. The analysis focuses on two distinct learning rate regimes: standard moderate step sizes and large step sizes that scale proportionally with the square root of the network width.
The investigation yields several key findings regarding model behavior. First, the article proves that the initial gradient update matrix is approximately rank-1 and aligns directly with the underlying target signal. Second, under a moderate learning rate, the model satisfies a Gaussian Equivalence property; the single gradient step consistently reduces prediction risk compared to the initial random features, yet the model remains confined to a linear regime and cannot defeat the best linear estimator on the raw input. Third, when the learning rate is scaled up proportionally to the square root of the network width, the neurons lose their near-orthogonality and break the linear barrier. In this large step-size regime, the updated features allow the model to achieve substantially lower prediction risk than standard fixed kernel lower bounds for certain nonlinear activation functions, achieving risk improvements that scale directly with the ratio of data dimensions to sample size.
These findings provide fundamental insights for practical machine learning engineering, training efficiency, and computational resource allocation. The results demonstrate that neural networks can realize substantial representational benefits immediately at the onset of optimization rather than requiring thousands of steps to escape the kernel regime. Crucially, the analysis reveals that learning rate magnitude during the initial phase acts as a structural switch: small learning rates restrict networks to linear approximations, whereas sufficiently large initial updates unlock true non-linear feature adaptation. This validates aggressive initial learning rate schedules and parameterization frameworks designed for large models.
For practitioners and engineering teams, the analysis suggests using sufficiently large learning rates during initial training phases to ensure the network exits the restrictive kernel regime and forms task-adapted features early. Organizations developing pretraining and transfer learning pipelines should prioritize capturing this initial phase of representation learning before fine-tuning readout layers. Furthermore, future technical investigations should explore intermediate learning rates to pinpoint the exact transition boundary between regimes, evaluate multi-step training dynamics, and extend the theoretical guarantees to settings where representation updates and regression are trained simultaneously on identical data samples.
Readers should interpret these findings within the scope of the article's foundational assumptions. The rigorous results rely on proportional asymptotic scaling, Gaussian input distributions, and single-index target functions, with feature updating and final readout regression performed on independent data splits. While the core qualitative behaviors are robustly supported by both rigorous mathematical proofs and empirical simulations, caution is warranted when extrapolating these specific quantitative error bounds directly to complex, non-Gaussian data architectures.
- Paper: Neural Tangent Kernel: Convergence and Generalization in Neural Networks, Arthur Jacot et al. (2018). This foundational paper establishes the Neural Tangent Kernel framework, providing the exact infinite-width, fixed-kernel baseline against which the source paper analyzes feature learning.
- Paper: Wide neural networks of any depth evolve as linear models under gradient descent, Jaehoon Lee et al. (2019). It characterizes the lazy, linear-evolution regime of wide neural networks under gradient descent, establishing the linear barriers that the source study aims to escape via initial gradient steps.
- Paper: Gradient Descent Provably Optimizes Over-parameterized Neural Networks, Simon S. Du et al. (2018). It proves the convergence of overparameterized two-layer networks by tracking Gram matrix dynamics near initialization, which directly informs the source's analysis of early-stage weight updates.
- Paper: Representation Learning: A Review and New Perspectives, Yoshua Bengio et al. (2012). It provides a foundational overview of representation and feature learning in neural networks, contrasting static features with data-adapted internal representations.
- Paper: Self-Consistent Dynamical Field Theory of Kernel Evolution in Wide Neural Networks, Blake Bordelon et al. (2022). It extends the mathematical study of feature learning beyond isolated gradient steps by establishing a dynamical field theory for continuous kernel evolution throughout training in wide neural networks.
- Paper: There Will Be a Scientific Theory of Deep Learning, Jamie Simon et al. (2026). It synthesizes asymptotic feature learning and lazy-versus-rich optimization dynamics into a broader, physics-inspired theoretical foundation for deep learning mechanics.
- Paper: Revenge of Monosemanticity: Specialized Neurons Improve Data Efficiency in MLPs, Amirhesam Abedsoltan et al. (2026). It investigates how individual neurons specialize and learn localized predictive directions across multi-cluster data, moving beyond single-index global representations.
- Paper: On the Role of Neural Collapse in Transfer Learning, Tomer Galanti et al. (2022). It applies representation learning principles to transfer learning by examining how geometric collapse in learned features governs generalization on unseen downstream tasks.
