arXiv Science⌕ Search

arXiv · 2609.34272

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.

Explore related subjects

Keep this discovery

Explore connections, maps & timelines

BibTeXRIS

Junlin Chen, Daize Dong, Huanwei Di, Haolong Jia, Jiawei Wu, Haotian Xie, Mingkai Zheng, Yang Li, Leshang Chen, Huishu Wang, Eric P. Xing, Hongyi Wang. 2026-09-28. Broken Symmetry in BF16 Attention: Why FlashAttention Gradients Blow Up Late in Training. https://arxiv.org/abs/2609.34272

Cite the original work for its findings. Save a collection to share your selection of sources.

KEEP EXPLORING

Related papers

Stochastic Engrams for Efficient Continual Learning

The ability to learn continuously in artificial neural networks (ANNs) is often limited by catastrophic forgetting, a phenomenon in which new knowledge becomes dominant. By taking mechanisms of memory encoding in neuroscience (i.e., engrams) as inspiration, we propose a novel approach that integrates stochastically-activated engrams as a gating mechanism for metaplastic binarized neural networks (mBNNs). This method leverages the computational efficiency of mBNNs combined with the robustness of probabilistic memory traces to mitigate forgetting and maintain the model's reliability. Previously validated metaplastic optimization techniques have been incorporated to further enhance synaptic stability. Compared to baseline binarized models and benchmark fully connected continual learning approaches, our method is the only strategy capable of achieving average accuracies over 70% in both class-incremental and domain-incremental MNIST benchmarks, matching full-precision state-of-the-art methods. Furthermore, we achieve a significant reduction in peak GPU and RAM usage, under 5% and 20%, respectively, as well as an ~8x reduction in memory footprint compared to full precision counterparts. Our findings demonstrate (A) an improved stability vs. plasticity trade-off, (B) reduced memory intensiveness, and (C) enhanced performance in binarized architectures. By uniting principles of neuroscience and efficient computing, we offer new insights into the design of scalable and robust deep learning systems.

cs.LG↗

DRAN: A Distribution and Relation Adaptive Network for Spatio-temporal Forecasting

Spatio-temporal forecasting remains challenging under non-stationary environments because both data distributions and spatial relations evolve over time. Temporal normalization and de-normalization are widely used to mitigate distribution shifts, but they may distort inter-node relationships and thereby impair spatial dependency modeling. To address these issues, we propose the Distribution and Relation Adaptive Network (DRAN) for spatio-temporal forecasting. DRAN incorporates a Spatial Factor Learner (SFL) module, which enables effective normalization and de-normalization while preserving spatial dependencies in spatio-temporal systems. To model evolving spatial interactions, DRAN further proposes the Dynamic-Static Fusion Learner (DSFL) module. DSFL decomposes features into static and dynamic components and adaptively fuses them according to input variability. Experiments on six benchmark datasets show that DRAN outperforms state-of-the-art baselines. Additional analyses demonstrate that SFL consistently reduces spatial-relation distortion across multiple normalization schemes, whereas DSFL captures complementary static and dynamic dependencies and adjusts their contributions according to temporal variability.

cs.LG↗

AYLA: Architecting a loss landscape in shallow neural networks to accelerate feature recovery

Feature learning in shallow neural networks exhibits rich yet fragile dynamics, including prolonged plateaus, abrupt phase transitions, and sensitivity to optimization hyperparameters. While recent theoretical work has characterized these behaviors through the geometry of loss landscapes, saddle escape mechanisms, and emergent scaling laws, practical methods for actively shaping these dynamics remain limited. In this paper, we introduce AYLA, a principled loss reparameterization framework that dynamically modulates gradient magnitudes during training without altering the location of stationary points or optimal solutions. AYLA applies a smooth, sigmoid-controlled power-law transformation to empirical loss, yielding a state-dependent effective learning rate that accelerates descent in flat or saddle-dominated regions while stabilizing late-stage optimization. Crucially, AYLA preserves all critical points of the original objective, acting solely as a monotone transformation that reshapes optimization trajectories rather than objectives. We evaluate AYLA in controlled teacher student settings using two-layer tanh networks trained on synthetic Gaussian data. Across stochastic gradient descent and multiple loss-exponent schedules, AYLA consistently improves feature recovery. This evidence is observed in terms of weight alignment, per-neuron cosine similarity, hidden-activation correlation, and spectral properties of learned representations, while AYLA maintains competitive or faster loss convergence. Spectral analyses further demonstrate that AYLA mitigates rank collapse and promotes richer internal representations, signaling a transition from lazy to active feature-learning regimes. AYLA offers a lightweight, theoretically grounded way to improve shallow-network optimization, especially in resource-limited or noise-sensitive settings.

cs.LG↗