cs.LGSep 28, 2026

Broken Symmetry in BF16 Attention: Why FlashAttention Gradients Blow Up Late in Training

Authors: Junlin Chen, Daize Dong, Huanwei Di, Haolong Jia, Jiawei Wu, Haotian Xie, Mingkai Zheng, Yang Li, +4 more

Organizations: Rutgers University · Carnegie Mellon University · Oracle · New York University · MBZUAI

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.

Figures & tables

Appendix figures & tables11 assets

Supplementary material from the paper’s appendix.

Appendix

Explore similar work

Aug 3, 2026cs.LG

One QK Channel, Many Sources: Tracing Low-Precision Attention Collapse

A bfloat16 transformer can train normally, then collapse abruptly. Prior work links collapse to structured attention errors and shows QK normalization disrupts their compounding. Distinct low-precision errors trigger the same collapse, leaving unclear whether each needs a fix at its source or one shared route can be blocked instead. We isolated the fault behind a reproduced GPT-2-class collapse to the streaming-softmax accumulator, where an fp32 streaming core repairs it, and turned it into an assay for moving a controlled error across sources. Using it, we found that errors placed outside attention still drove the same QK spectral runaway, and that correcting only QK kept training stable while the fault stayed active. This is a source-channel dissociation: fault source is not failure channel. It held across tested architectures and scales, and reproduced on a second GPU architecture. As a causal probe, projecting each update off the current QK weights' leading three singular directions held the query projection's largest singular value to 11.1, whereas removing equal energy elsewhere left it at 237: the QK channel causally drives the early runaway. What lets the injected error in is temporal sign-coherence, its per-head sign persisting across steps, not aggregate deviation; once inside, the runaway shows as attention-logit saturation. QK-Guard, a dormant controller, tests this by switching on parameter-free QK normalization at the first monitored threshold crossing. On the runs designated for this test, the QK-local action prevented the failure of each matched or same-configuration unguarded run; on plain GPT-2, all 12 final train and validation losses were within 0.03 nat of same-configuration always-on QK-norm, and both methods ran 60k steps without collapse. Intervention at the QK locus therefore suffices in place of a fix at each source.
Sep 21, 2026cs.LG

FlashBoB: I/O-Efficient Exact Backward-over-Backward for Softmax Attention

Transformer models built on the attention mechanism have become a central building block in modern deep learning, yet softmax attention remains a major bottleneck for long-context workloads. While FlashAttention makes the forward and first backward passes I/O-efficient, it does not support backward-over-backward (BoB), which enables exact differentiation through the backward pass for applications such as second-order optimization, test-time training, gradient-based memory, and meta-learning. Existing BoB implementations either materialize large intermediate tensors or exhaust GPU memory at long sequence lengths. We present FlashBoB, an exact, I/O-efficient algorithm for BoB in softmax attention that keeps computation within on-chip tiles and avoids all N×NN \times N intermediate tensors, where NN is the sequence length. The key insight is a hierarchical affine structure in the softmax double backward: two row-wise scalars determine all outputs through affine transformations. This yields a two-pass schedule with bounded on-chip static random-access memory (SRAM) usage and minimal off-chip high-bandwidth memory (HBM) traffic. FlashBoB achieves Θ(N2d2/M)Θ(N^2 d^2/M) HBM traffic (dd is the head dimension and MM is the memory size) and, within the standard FlashAttention-style score-recomputation model, matches the inherited large-cache lower bound for exact forward attention. Empirically, it scales exact attention BoB to N=262KN=262\text{K} on a single A100 80GB GPU, where prior PyTorch exact baselines fail by N=16KN=16\text{K}, and is up to 6.3×6.3\times faster than FlashBack. These results make exact second-order attention practical at long-context sequence lengths where prior implementations cannot run efficiently.
Jun 22, 2026cs.LG

FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training

Reverse-mode differentiation computes every weight gradient, writes it to memory, and only then lets the optimizer read it back. This two-phase schedule sets the memory ceiling of modern training: at the seam between the phases, every layer's gradient is live at once. We argue that this materialized gradient is an artifact of how differentiation is staged, not a quantity that learning requires -- and we eliminate it. FORGE folds the optimizer step into the backward pass and applies it one tile at a time, entirely in registers, so each gradient tile is consumed the instant it is produced and never becomes a tensor. The fusion changes only when the update happens, not what it computes: in full precision the fused step is provably exact -- the identical optimizer update, for every element-wise rule -- and that exactness survives tensor- and sequence-parallel sharding; in the bf16 and 8-bit regimes used in practice it is faithful rather than bit-identical, its deviation bounded and, for the weight store, rendered unbiased by stochastic rounding. Because each gradient tile is born and consumed in the same registers, it is never converted down to bf16 to be stored and read back; FORGE thus preserves the full-precision fidelity that both bf16 and 8-bit optimizers lose to that conversion. Nor is the method tied to one architecture or one optimizer: linear layers are ubiquitous, and FORGE reclaims the gradient memory of any of them under any element-wise rule. Empirically FORGE more than halves the memory of an optimizer step and, at the small batch sizes typical of fine-tuning and continued pretraining, runs about 1.5x faster; integrated into tensor-parallel Megatron-LM it fits 8B training at four times the micro-batch a standard optimizer allows on the same GPUs.