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

This presentation examines a subtle but catastrophic numerical failure in BF16 fused attention backward passes that appears only late in training. We explore how two distinct error channels—one in the saved attention output and one in the score-gradient casting—violate a fundamental zero-sum conservation law, causing gradients to explode by three orders of magnitude without triggering conventional overflow alarms. The talk demonstrates how a simple projection repair restores the broken symmetry and matches FP32 accuracy at only 4.7 percent overhead.
Script
Train a 450 million parameter transformer with BF16 fused attention for 50 billion tokens, and somewhere around 25 billion tokens, the gradient norm silently explodes by three orders of magnitude. No NaNs, no infinities, just a sudden 0.2 nat loss drift that looks like an optimizer bug but is actually a hidden numerical defect in the attention backward pass.
The authors traced the failure to just two attention layers. When they replaced only the backward computation of layers 5 and 11 with FP32 arithmetic while keeping everything else bitwise identical, the median gradient norm dropped from 5353 to 18. The pathology was not a forward loss anomaly or an optimizer failure, it was a backward numerical error localized to specific layers.
Here is the core defect. Exact attention backward has a beautiful conservation law: the score gradient must sum to zero across each row, which makes query gradients invariant to shifting all keys by the same offset. But BF16 casting breaks that invariant. The rounded score gradient no longer sums to zero, and the nonzero residual mass gets multiplied by the mean key, creating a spurious gradient component that grows with key offset rather than key spread.
The projection fix is conceptually simple. After casting the score gradient to BF16, subtract its row sum weighted by the attention probabilities, restoring exact zero mass. This removes the mean key leakage term while preserving all other BF16 rounding behavior. Median query gradient error drops from 773 percent to 0.34 percent, matching fused FP32 attention.
In matched 50 billion token runs, GProj and FP32 attention finish at essentially identical loss, 1.665. The native FA3 run diverges to 1.866. Fixing only the saved output keeps gradients from exploding but still drifts to 1.679 and drives layers into the same extreme logit regime. The lesson is sharp: a low precision kernel must respect exact structural identities, not just minimize aggregate numerical error.
GProj restores the broken symmetry at 4.7 percent training step overhead and eliminates a failure mode that conventional overflow monitoring cannot detect. If you want to explore this paper further or create your own video summaries of cutting edge research, visit EmergentMind.com.