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

📅 2026-09-28
📈 Citations: 0
✨ Influential: 0
📄 PDF
🤖 AI Summary
This study addresses the instability of FlashAttention under BF16 precision, where rounding errors violate the zero-sum constraint of softmax, inducing symmetry breaking and query gradient leakage that ultimately triggers gradient explosion during late-stage training. To resolve this, we elucidate the underlying mechanism and propose GProj, a canonical projection method that restores the zero-sum constraint via rank-one correction. Furthermore, we introduce an efficient correction framework integrating FMA-based forward repair with high-precision backward recomputation. Our approach reduces gradient errors to FP32-comparable levels (<0.4%) while incurring only 4.7% additional step-time overhead, thereby ensuring convergence stability for large-scale pretraining at minimal computational cost.
📝 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.
Problem

Research questions and friction points this paper is trying to address.

BF16
FlashAttention
gradient blow-up
broken symmetry
softmax
Innovation

Methods, ideas, or system contributions that make the work stand out.

BF16 Attention
Broken Symmetry
Gradient Explosion
Gauge Projection (GProj)
FlashAttention
🔎 Similar Papers