Mind the Drift: Diagonal Linear Networks Under Large Learning Rates
Authors: Aniket Sanyal, Tom Jacobs, Rebekka Burkholz
Organizations: CISPA Helmholtz Center for Information Security · Technical University of Munich · Konrad Zuse School of Excellence in Learning and Intelligent Systems (ELIZA)
Large learning rates can qualitatively change the trajectory of neural network training, often pushing optimization into regimes far from classical gradient-flow behavior. The Edge of Stability (EoS) offers a valuable lens on the dynamics such learning rates induce. We study corresponding dynamics in diagonal linear networks, where we uncover a competition between two distinct implicit biases that jointly determine the sparsity of the recovered solution in regression settings. Complementary to the Gain, which captures the average discretization error accumulated by Gradient Descent relative to Gradient Flow, we derive a closely associated but overlooked quantity: the Drift. Under large learning rates, it describes an imbalance between different discretization errors and represents a systematic shift in the optimization trajectory. While the Gain grows monotonically in certain regimes, and can bias towards denser, flatter interpolators, the impact of the Drift depends on its alignment with potential solutions, which can either counteract or reinforce the effect of the Gain. Consequently, its behavior drives model selection, particularly during early training epochs. To validate our theoretical insights, we introduce an intervention that actively steers the Gain to recover sharper, sparser solutions. Thus, our analysis reveals that large learning rates do not universally hinder the recovery of sparse solutions. On the contrary, they can be harnessed to control the implicit bias of training.
Figures & tables
Figure 1: A diagonal linear network with learning rate γ=0.3 and 0.4 to show the evolution of the predictor coefficient. The smaller step size does not reach the Edge of Stability regime, while the larger one does. We recover a sparser solution for the former.
Figure 2: Alternating behavior of the Gain and the Drift. At initialization, the training starts from the left and proceeds to the right as the gradients shrink. The upper half of the plot indicates a sparsity promoting behavior, while the bottom half represents bias towards denser predictors. The two dotted lines represent the support and offsupport poles. See Table 1 for a summary.
Figure 3: A diagonal linear network example case with control parameter B=3.1 and learning rates γ=0.3 and 0.4 . ( Left ) We plot the product of sharpness and step size. The smaller step size run converges below the EoS regime while larger learning rate oscillates. ( Middle ) In both cases, the Gain and the Drift accumulate comparatively (Lemma 2.5 and 2.6 ). ( Right ) We plot the predictor coefficient, recovering a sparser solution for the former run and denser for the latter.
Figure 4: Sparse recovery for GD, SGD and GC on the noiseless sparse regression problem. ( Left ) Distance to the sparse teacher. Due to its implicit bias (Prop. 3.3 ), GC recovers a predictor substantially closer to the sparse teacher. ( Right ) Loss curve for the three methods. The learning rate switch is to allow GC to converge.
Figure 5: Accumulated Drift and Gain of all Query-Key pairs for a DeiT-Small trained on ImageNet with Muon, across three (batch size, learning rate) configurations. Both quantities grow substantially early in training, but Drift starts plateauing while Gain continues to accumulate throughout training.
Appendix figures & tables6 assets
Supplementary material from the paper’s appendix.
Appendix
Figure 6: Sparse recovery in noiseless regression with γ=0.1 . Left: normalized recovery error. Right: normalized training loss. GC achieves substantially lower recovery error although it converges slower compared to (S)GD.
Figure 7: Sparse recovery in noiseless regression with cosine annealing from γmax=0.35 to γmin=0.05 . Left: normalized distance to the sparse teacher. Right: normalized training loss; GC recovers the sparse teacher substantially better than (S)GD.
Figure 8: The Drift is necessary for sparse recovery. At every checkpoint k we minimise the implicit-bias potential over the interpolation manifold M , using the Gain Qk and Drift Dk accumulated so far, and plot the distance of that minimizer to the sparse teacher. With the Drift, ψak+21⟨Dk,⋅⟩ (dashed) converges exactly to the GD iterate (blue), as predicted by Prop. 2.3 . Keeping the same Gain but setting D=0 (purple) selects a much denser interpolator. The gradient-flow potential ψα2 (dotted, Q=D=0 ) is shown for reference. GF recovers the sparsest solution. The shaded region marks the oscillatory phase ( γkSk≥2 )
Figure 9: Following the onset of EoS oscillations, Gain continues to accumulate while the net change in Drift remains small. Measuring both quantities relative to Kedge makes this separation explicit, supporting the observed saturation of Drift during the oscillatory phase.
Gain Bias
Drift Bias
Regime
Direction
Solution
Direction
Solution
Above both poles
−h
βs∗
+h
βd∗
Between the poles
−h→+h
-
+h→−h
-
Below both poles
+h
βd∗
−h
βs∗
Appendix
Table 1: Directional biases of Gain and Drift toward the sparse solution βs∗ and dense solution βd∗ . Between the poles, arrows indicate alignment reversals as the residual magnitude decreases.
Batch size
Learning rate
QK effective rank
Val. accuracy
128
0.02
1959.05±38.49
79.68±0.096
128
0.04
1815.71±19.38
79.24±0.161
256
0.04
1980.98±31.96
79.56±0.064
Appendix
Table 2: Effective rank and validation accuracy on ImageNet. Error bars denote 95% confidence intervals over three seeds. The maximum effective rank possible is 4608 .
We study the gradient flow dynamics of diagonal linear networks for regression tasks under infinitesimal initialization. Extending Theorem 1 from Pesme & Flammarion (2023), we generalize the analysis to both deep diagonal linear networks and a broader class of two-layer diagonal linear networks (as defined in Definition 4.1). Specifically, we demonstrate that the training trajectories of these models can be equivalently characterized by the proposed Algorithm 1. We further prove that this algorithm converges to the solution of a modified l1 norm minimization problem. As a result, we establish that the implicit bias of both network architectures corresponds to a modified l1 norm in the regime of infinitesimal initialization. Additionally, we provide insights into the underlying mechanisms governing these dynamics by identifying the Structural Invariant Manifold (SIM) (Zhao et al., 2026) as the key geometric structure that shapes the learning process.
Jiajie Zhao, Jianxing Wang, Junjie Yang +2
School of Mathematical Sciences, Shanghai Jiao Tong University · Institute of Natural Sciences, MOE-LSC, Shanghai Jiao Tong University · School of Artificial Intelligence, Shanghai Jiao Tong University
We study optimal learning-rate selection in two-layer and three-layer linear neural networks trained to learn linear target functions. In particular, we derive the exact closed-form expressions for the gradients and test loss after one and two steps of gradient descent, enabling a precise characterization of early training dynamics. We characterize how learning rates should scale under the gradient approximation in the first two steps, and prove that performing updates with this approximation yields a tractable surrogate loss with a tight, small approximation error. This formulation enables the theoretical analysis of layer-wise learning rates and reveals a distinct early-training regime: test loss can be minimized by unequal learning rates at the initial step, while equal learning rates become optimal in subsequent steps. Our numerical experiments validate the theory and demonstrate the importance of balancing layer-wise learning rates early during training. The code is available at: https://github.com/TDCSZ327/Layer-Balancing.
Tianyu Pang, Vignesh Kothapalli, Shenyang Deng +3
Dartmouth College · Stanford University · Virginia Tech
We study the dynamics of gradient descent in the Edge of Stability regime, where the learning rate is large enough to induce persistent oscillations in the trajectory, which has been linked to better generalization performance. We introduce the mean--fluctuation dynamics, a tractable continuous-time model coupling the window-averaged trajectory to its fluctuation covariance. Among our contributions, we rigorously derive this model from gradient descent in a sharp-valley framework, characterize its stationary states and their linear stability, and establish precise connections with other effective dynamics. Numerical experiments illustrate these predictions and their finite-time limitations. We also study our model in the overparametrized regime of wide two-layer networks at a fixed learning rate, where we rigorously derive a kinetic equation describing weights and their fluctuations as a Wasserstein-2 gradient flow, for which we prove well-posedness, a mean-field limit, and conditional convergence results.
Antonin Chodron de Courcel
Ecole Normale Sup´erieure, CNRS, 45 rue d’Ulm, 75005 Paris, France