Papers for
deep learning framework developers
Papers whose findings have a practical use for this group, as judged from the abstract. Open a paper to read what it means in practice.
Bfloat16 errors cause large transformer training instabilities fixed by gauge projection
Broken Symmetry in BF16 Attention: Why FlashAttention Gradients Blow Up Late in Training
Abstract: BF16 is now standard in large-scale pretraining, including in fused attention kernels such as FlashAttention, and these kernels are widely trusted. When we used FlashAttention-3 to pretrain a 450M-parameter transformer on 50B tokens, however, we ran into a problem: training was healthy for 25B tokens, then the gradient norm grew a thousandfold and the loss ended 0.2 nats above FP32 attention, without a single NaN. Recomputing the attention backward of just two layers in FP32 removes almost all of the excess gradient. Part of the cause is known: a fused multiply-add in the forward softmax, so far treated as an extreme-input NaN case and never fixed in FlashAttention-3. Repairing it stops the blow-up, but the query gradient is still wrong by more than its own size, and training still drives attention logits to thousands of times their size under accurate gradients. The remaining error comes from a broken conservation law. The softmax score gradient sums to zero along every row, which makes the query gradient blind to where the keys sit as a group; rounding it to BF16 leaves a small nonzero sum that leaks the mean key into the gradient, and the leak grows exactly as late training makes keys large and attention sharp. We introduce GProj (gauge projection), which restores the zero sum after the cast with two rank-one corrections per row. It cuts the remaining median query/key gradient errors from 219%/13% to 0.34%/0.37%, on par with FP32 attention, for 4.7% more time per training step. In matched from-scratch runs it trains to the same loss as FP32 attention, while FlashAttention-3 and key smoothing both destabilize.
Adaptive spectral estimation improves matrix optimizers in ai training
Cost-free Spectral Estimation for Adaptive Newton--Schulz in Matrix Optimizers
Abstract: Matrix optimizers such as Muon transform each momentum matrix through an approximate orthogonalization, typically implemented by a small number of Newton-Schulz matrix multiplications. The quality and cost of this approximation depend strongly on the singular-value spectrum of its input, yet existing implementations use the same fixed polynomial routine for every layer and throughout training. We show that this uniform treatment is unnecessary: the computations in the Newton-Schulz method already reveal enough information to make the method adaptive. The Gram matrices formed inside Newton-Schulz iterations yield spectral moments through inexpensive scalar reductions, requiring no additional matrix multiplications. From these moments, we recover an estimate of the empirical singular-value distribution and use it to select a polynomial routine specialized to the current matrix. This turns Newton--Schulz orthogonalization into a spectrum-adaptive procedure that responds to differences across both layers and training time. On saved momentum matrices, spectral estimation substantially reduces orthogonalization error at a fixed iteration budget or reaches the same accuracy with fewer iterations, and in GPT pretraining up to 1B parameters it lowers the validation loss of two matrix optimizers. Our results suggest that matrix-function operations inside optimizers need not be designed for a conservative worst-case spectrum: they can cheaply measure the spectrum they are already processing and specialize computations accordingly.
Momentum helps stabilize high learning rates for faster training progress
Towards Understanding Momentum Acceleration in River-Valley Loss Landscape
Abstract: The empirical success of pretraining large language models has inspired a deeper investigation into the underlying loss landscapes and the optimization dynamics. Recent empirical and theoretical study suggest that the training loss landscape often exhibits a "river-valley" structure, which features a low-loss manifold (river) flanked by sharp orthogonal directions with higher loss (mountains). In the long term, the optimization progress is determined primarily by the progress along the river. Within such a landscape, gradient descent with large learning rates can move faster along the river despite high apparent loss due to vertical oscillations, while a subsequent sharp decay in the learning rate suppresses these oscillations, revealing genuine optimization progress. This explains the recent success of warmup-stable-decay (WSD) learning rate scheduler which, unlike cosine scheduling, keeps stable high learning rate and decays before producing intermediate checkpoints. Building on this foundation, in this work we take a step further and study the role of momentum within such a loss landscape. We establish theoretical analysis that characterizes how momentum accelerates optimization by stabilizing large learning rates that can not be tolerated by vanilla GD without deviating significantly from the river. The enabled large learning rate in-turn gives greater speed along the river and makes faster essential progress in the long run. Another intriguing observation from theory is that for a river-valley landscape with very flat and slow-spinning river, the momentum itself does not contribute directly to acceleration in terms of the speed of tracking the river, while the main acceleration comes from the admissible larger learning rate.
Cosine relations reshape momentum for better training outcomes
COREM: Cosine-Relation Momentum Reshaping with Stateful Writeback
Abstract: Matrix-valued optimizer states may contain relational structure that is not captured by treating their entries independently. We study whether relations within matrix-valued optimizer states can be exploited to improve optimization. To this end, we introduce a unit-relation-transform abstraction and instantiate it as COREM, a Cosine-Relation Momentum Reshaping method with stateful writeback. COREM partitions the momentum state into update units, computes cosine relations among them, and uses these relations to reshape the momentum before writing the transformed state back to the optimizer. This stateful mechanism allows the reshaped momentum to affect not only the current update but also future optimization dynamics. We evaluate COREM on CIFAR-10 with an MLP and on enwik8 with a Transformer. Compared with Muon, COREM shows lower early-stage step efficiency but stronger improvement in the mid-to-late stages of training, achieving better final validation performance on CIFAR-10 and comparable final performance on enwik8. Spectral diagnostics on enwik8 show that COREM consistently increases entropy effective rank and reduces the concentration of singular energy in dominant modes, while preserving an anisotropic spectrum. For square matrix updates, COREM requires approximately 13.3% of the transformation FLOPs of Muon with five Newton-Schulz iterations.
Scale invariant optimization shows sharp stability boundary with weight decay
When Does Scale-Invariant Optimization Become Unstable? An Exact Schedule Law with Weight Decay
Abstract: Normalization renders large parts of neural networks effectively scale invariant, inducing a hidden feedback loop in which learning-rate schedules and weight decay interact through the parameter norm to control the effective step taken by the optimizer. We show that this interaction is governed by an exact discrete-time law: a single scalar quantity captures all schedule and decay forcing, while norm growth induces an opposing geometric self-quenching effect. This yields a sharp boundary that cleanly separates contraction- and expansion-dominated effective learning rate regimes. To understand the underlying mechanism, we provide exact analysis of a fully solved normalized regression model where the dynamics reduce to two dimensions and show that the balance point is intrinsically unstable, implying that constant learning rate with weight decay cannot stably maintain an interior equilibrium and instead produces recurrent behavior driven by discrete-time Jacobian structure. We further extend this perspective across optimizers through unified homogeneous-optimizer framework that reveals a structural dichotomy in self-quenching strength, providing a first-principles explanation for why adaptive methods exhibit systematically weaker stabilization under normalization. Across dynamical systems and neural networks (MLP, CNN, GPT2 / MNIST, CIFAR, wikiText, OpenWebText), the predicted law holds with high precision and enables direct control of training via the identified scalar, with performance peaking sharply at the predicted boundary. Together, these results isolate a single governing quantity for scale-invariant optimization, providing a precise and actionable lens on training dynamics, optimizer behavior, and schedule design in modern deep learning. Code is available in https://github.com/shasanamin/normalized-optimization-dynamics.
Neural networks generalize differently across layers with weight decay training
A Theoretical Analysis of Generalization Dynamics in Neural Networks under Gradient Descent with Weight Decay
Abstract: Understanding generalization remains a central challenge in machine learning because it requires jointly considering data, architecture, and training dynamics. In this paper, we develop a theoretical framework that characterizes how these factors jointly shape generalization performance throughout training. More precisely, we study a broad class of neural networks trained under the $\ell^2$ loss by gradient descent (GD) with weight decay, and prove the convergence of GD to a neighbourhood of the global minimizers of the empirical loss. By partitioning the space based on the input data, we then decompose the population error into data error, optimization error, and prediction variation error, and bound them separately. In particular, for the prediction variation error, which measures the oscillations of the learned function, we propose (local) approximate homogeneity and derive explicit cellwise and layerwise bounds for its evolution along the training trajectory. These bounds yield two important implications: a necessary condition of improved generalization explains differences in layerwise generalization behavior; a sufficient condition describes delayed generalization and provides a theoretical characterization of grokking.