Papers
Topics
Authors
Recent
Search
2000 character limit reached

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

Published 28 Sep 2026 in cs.LG, cs.DC, and math.NA | (2609.34272v1)

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.

Summary

  • The paper identifies two numerical issues in BF16 fused attention backward passes: a saved-output error due to fused softmax arithmetic and a conservation-law violation after BF16 casting, both leading to large gradient errors.
  • The authors propose a corrected version of FlashAttention-3 (FA3), called GProj, which reduces the median gradient error to approximately 0.34% and 0.37% for query and key gradients, respectively, preserving computational efficiency and reducing large-logit regime issues.
  • Experimental validation shows that GProj matches the performance of FP32 attention in stability and loss, resolving drift and incorrect gradient growth in both synthetic and practical training contexts, ensuring reliable training without NaNs or overflows.

Problem setting and empirical failure

The paper investigates a silent numerical failure in BF16 fused attention backward passes. The motivating experiment is a 450M-parameter transformer trained for 50B tokens with FlashAttention-3 (FA3). Training appears healthy for approximately 25B tokens, after which the pre-clipping gradient norm increases by roughly three orders of magnitude and the loss diverges from an otherwise matched FP32-attention run by approximately 0.2 nats. No NaNs or infinities occur. The failure is therefore not detected by conventional overflow monitoring and could plausibly be misattributed to optimization, data, or model-scale effects.

The authors localize the excess gradient to the attention backward. At 33.6B tokens, replacing only the backward computations of layers 5 and 11 with FP32 arithmetic, while holding all forward activations bitwise fixed, reduces the full-model gradient norm from a median of 5353.856 to 18.185. Replacing the saved attention output alone reduces it to 19.746. This intervention establishes that the dominant pathology is a numerical backward error in two attention layers rather than a forward loss anomaly or a general optimizer failure.

The training trajectory exhibits a pronounced delayed-onset behavior. FA3 remains competitive through the early and middle stages, then enters a regime in which layers 5 and 11 develop exceptionally large query and key norms, nearly one-hot attention distributions, and increasingly inaccurate gradients. The delayed onset is important: random well-conditioned validation inputs and early-training comparisons do not expose the defect.

Figure 1

Figure 1: Matched from-scratch training shows late gradient growth and loss drift for FA3, whereas GProj remains aligned with FP32 attention throughout 50B tokens.

The paper identifies two distinct numerical channels. The first is a known forward-softmax issue: FA3 fuses scaling and maximum subtraction in an FMA, so the maximum score need not map exactly to zero. The resulting error contaminates the BF16 attention output saved for the backward pass. The second channel is the paper’s principal contribution: even with an exact saved output, casting the score gradient to BF16 violates an exact zero-sum invariant and produces a spurious gradient component proportional to the mean key.

The saved-output channel

For a query row, keys kjk_j, values vjv_j, attention probabilities pjp_j, and incoming output derivative uu, the exact score-gradient entries are

gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),

where aj=u⊤vja_j=u^\top v_j and μa=∑jpjaj\mu_a=\sum_jp_ja_j. Consequently,

∑jgj=0.\sum_j g_j=0.

The query gradient can be written as

GQ=α∑jgjkj,G_Q=\alpha\sum_j g_j k_j,

with α\alpha denoting the attention scale. FA3 estimates the scalar reduction vjv_j0 through the BF16 output saved during the forward pass. If the saved output induces an error vjv_j1, the resulting query-gradient error is

vjv_j2

where vjv_j3 is the attention-weighted mean key. Thus, even a small saved-output error is amplified when the mean key is large.

The source of the saved-output error is traced to the fused softmax arithmetic. The implementation effectively computes a fused scaled-and-shifted exponent argument rather than explicitly subtracting the row maximum before scaling. Because the FMA does not necessarily round the intermediate product separately, the maximum element may receive a nonzero exponent argument. In a one-key row, where exact attention requires vjv_j4 and vjv_j5, this can produce vjv_j6 and a nonzero query gradient.

The paper verifies this mechanism using 96 prescribed configurations whose dot products are analytically controlled. A scalar model of the FMA, BF16 exponentiation operand, and FP32 denominator reproduces the native outputs bitwise, including all six cases in which the one-key identity fails. The proposed subtract-before-scale modification, FA3-SBS, restores the identity in all prescribed cases and substantially reduces the observed gradient error.

The saved-output channel accounts for most of the initial failure. Across a confirmation set of 127 documents, the predicted saved-output error has near-perfect alignment with the observed native query-gradient error in the most affected layers. At checkpoint 8000 in layer 11, replacing only the saved output reduces the absolute query-gradient error by a median factor of 72.02. However, this repair does not eliminate the remaining backward defect.

Figure 2

Figure 2: Across 127 documents, the saved-output prediction explains most of the native query-gradient error, with high error cosine and a small unexplained residual.

The distinction between saved output and saved log-sum-exp is established by state-cross experiments. Substituting the corrected output while retaining the native log-sum-exp reproduces the improvement, whereas substituting only the log-sum-exp does not. The implication is direct: the forward repair must target the saved attention output, not merely normalization statistics.

The broken conservation law

The central theoretical observation is that exact attention backward propagation possesses a translation symmetry. Adding a common vector vjv_j7 to every key in a row changes every logit by the same constant, leaves the softmax probabilities unchanged, and therefore leaves the exact local VJP unchanged. Algebraically, this invariance follows from the zero row sum of the score gradient:

vjv_j8

The query gradient can equivalently be expressed in centered form,

vjv_j9

Only key differences matter. A numerical backward that responds to the absolute common offset of the keys violates this symmetry.

FA3 casts the score gradient pjp_j0 to BF16 before contracting it with the keys. Let pjp_j1 be the cast operand and define the rounding error pjp_j2 and row mass pjp_j3. Since the exact score gradient has zero mass,

pjp_j4

The first term is an ordinary centered rounding error. The second is the leakage term. It is generated solely by the nonzero row sum left by BF16 rounding and is multiplied by the mean key. A common key translation pjp_j5 leaves the exact gradient unchanged but changes the numerical contraction by pjp_j6.

This mechanism becomes severe in the late-training regime. When attention is nearly one-hot, the exact query gradient is a covariance over a distribution with very small effective spread and can be close to zero. The mean key, however, can remain large. A single BF16 rounding unit in the score gradient can therefore dominate the true signal. The paper’s explicit four-key witness makes the effect non-asymptotic: with an exact forward and exactly representable inputs, BF16 casting changes an exact zero query-gradient coordinate into a value of pjp_j7.

Figure 3

Figure 3: BF16 casting leaves nonzero score-gradient row mass, producing a mean-key leak whose magnitude grows with key offset; GProj removes the leak while retaining BF16-level error.

This result is stronger than a diagnosis of inaccurate softmax outputs. It shows that a perfect forward pass cannot remove the principal residual error, because the defect is introduced after the forward state has been computed. It also explains why FP32 accumulation alone is insufficient: the incorrect BF16 multiplicands have already broken the invariant before the matrix products are accumulated.

The theoretical distinction between forward repair and projection is formalized through a key-translation argument. A backward pass that directly contracts the cast score gradient has an error component proportional to pjp_j8, and its worst-case relative error is unbounded under common key translations. A projected backward pass removes this dependence and leaves an error bounded by the attention-weighted key spread. The result identifies the relevant numerical quantity not as the absolute key magnitude but as the key offset relative to within-row variation.

Gauge projection

GProj restores the zero row sum after the score-gradient cast. The implementation uses the BF16 probability operand pjp_j9 that the kernel already forms and its actual mass

uu0

Given uu1, GProj subtracts uu2:

uu3

This correction satisfies uu4 even when the BF16 probability mass is not exactly one. Subtracting uu5 without dividing by uu6 would leave a residual mass uu7, which is why the correction must use the mass of the actual operand consumed by the kernel.

The corrected query contraction is implemented without materializing uu8:

uu9

The key gradient requires a second pass because gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),0 is known only after a query row has processed all key tiles, whereas the native key accumulator is organized by key tile. GProj therefore adds a second contraction to correct gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),1 before its final BF16 cast. The value gradient remains native.

The projection is exact on any zero-sum score gradient and acts only on the cast error. In geometric terms, it removes the component of the numerical error that lies in the row-mass direction while preserving the other BF16 rounding components. The corrected query contraction depends on centered keys and is invariant to common key translation. Its remaining error is consequently controlled by key spread rather than key offset.

Gradient accuracy and training behavior

The principal accuracy evaluation uses eight captured attention inputs from layers 5 and 11 at 23.1B and 33.6B tokens. All methods receive identical BF16 inputs, and gradients are compared with an FP64 centered VJP. FA3-SBS and GProj have byte-identical forward outputs, so their comparison isolates the backward.

Method gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),2 error gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),3 error gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),4 error
FA3 773% 85.6% 0.369%
FA3-SBS 219% 13.3% 0.334%
GProj-Q 0.342% 13.3% 0.334%
GProj 0.342% 0.371% 0.334%
Fused FP32 attention 0.385% 0.371% 0.165%

The forward repair reduces the median gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),5 error from 773% to 219%, but leaves the query error larger than the gradient itself. The query projection reduces it further to 0.342%, essentially matching fused FP32 attention. The additional key pass reduces gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),6 error from 13.3% to 0.371%. These results separate the roles of the two corrections: GProj-Q repairs the query contraction, while the second pass repairs the key contraction.

The synthetic suite contains 58 stress cases covering common-key translations, singleton rows, packed layouts, unequal lengths, tile boundaries, sharp attention, tiny upstream derivatives, and long sequences. GProj passes all cases and does not increase gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),7 or gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),8 error relative to FA3-SBS. Its largest reported synthetic relative errors are 2.49% for gj=pj(aj−μa),g_j = p_j(a_j-\mu_a),9 and 2.76% for aj=u⊤vja_j=u^\top v_j0, both occurring in numerically small-gradient cases. For exact-zero targets, residuals remain near the arithmetic floor; for example, singleton-row aj=u⊤vja_j=u^\top v_j1 is exactly zero and the corresponding aj=u⊤vja_j=u^\top v_j2 residual is approximately aj=u⊤vja_j=u^\top v_j3.

The training study uses matched from-scratch runs with the same seed, data order, architecture, and optimization configuration. GProj and FP32 attention remain stable through 50.0B tokens and finish with essentially identical loss: 1.665. FA3 finishes at 1.866 over the final 0.84B tokens, approximately 0.20 nats worse, with a median gradient norm of 30.705. FA3-SBS remains numerically stable in the narrow sense that its gradient norm does not exceed 10, but its final loss is 1.679 and it drives several layers into the same extreme-logit regime. Key smoothing delays the failure by approximately 3B tokens but ultimately produces a final loss of 1.903 and a median gradient norm of 233.384.

Figure 4

Figure 4: Holding the forward fixed localizes the excess full-model gradient to the backward of layers 5 and 11; correcting the saved output removes most of the error, while FP32 removes both numerical channels.

The regime audit clarifies why stable loss alone is insufficient. At 33.6B tokens, FA3 has logit scales of 12,843 and 51,057 in layers 5 and 11, respectively; 99.4% of layer-11 rows satisfy aj=u⊤vja_j=u^\top v_j4. GProj has maximum logit scale 26, layer-11 query/key RMS values of 2 and 3, and no rows with aj=u⊤vja_j=u^\top v_j5. Thus, the inaccurate backward does not merely perturb an otherwise healthy trajectory: it drives the model toward a large-key, nearly one-hot regime that amplifies the very error responsible for the instability.

The released Qiu–Yao dynamic-shift baseline performs worse in this setup. Its shift rule can move positive near-tied tile maxima by aj=u⊤vja_j=u^\top v_j6, making all tile weights exponentially small when logits are large. Tile sums then fall below the implementation’s aj=u⊤vja_j=u^\top v_j7 clamp. The run spikes at approximately 6.3B tokens and is stopped at 13.2B tokens. This result does not refute dynamic shifting generally, but it shows that the particular released arithmetic and clamp policy do not provide a reliable substitute for conservation-aware backward correction in this evaluation.

Computational cost

GProj preserves the BF16 matrix products and does not require additional forward state or quadratic attention storage. On an H200 with batch one and sequence length 4096, complete training-step time is 238.120 ms for GProj versus 227.357 ms for FA3, a 4.73% increase. Eager FP32 attention requires 719.420 ms, or 216.43% additional time. An optimized fused FP32 implementation still costs 308.276 ms, 33.43% above FA3.

Peak allocated memory is effectively unchanged: 9.191 GiB for both FA3 and GProj in the reported comparison. The overhead is therefore principally computational, arising from the additional aj=u⊤vja_j=u^\top v_j8 contraction, row-statistic handling, and second key-correction pass. The reported 4.7% figure is specific to the complete training step and the tested 4096-token configuration; it should not be interpreted as a sequence-length-independent kernel overhead.

Limitations and open questions

The empirical scope is narrow. The main experiments use BF16 FA3 on Hopper GPUs, head dimension 64, grouped-query attention, sequence length 4096, and a single 450M-parameter architecture. Although the fixed-operand backend sweep shows related backward errors in FA2, cuDNN, and PyTorch fused SDPA, the paper does not establish equivalent late-training behavior across GPU generations, head dimensions, precision formats, model families, or training scales.

GProj does not correct aj=u⊤vja_j=u^\top v_j9, whose error remains that of FA3-SBS. Its projection also assumes that the relevant conservation law is the appropriate invariant for the implemented attention formulation and that the BF16 probability mass remains finite and nonzero. The synthetic evaluation satisfies these conditions, but more adversarial kernels or alternate sparsity and quantization schemes may require different projections.

Finally, the training evidence consists of matched runs with one principal seed and a fixed model/data configuration. The agreement between GProj and FP32 attention is strong within that setting, but the extent to which the same mechanism controls instability in larger models, longer contexts, FP8 kernels, or other fused backward implementations remains an open empirical question.

Conclusion

The paper shows that BF16 fused attention can produce severely incorrect gradients without NaNs or obvious forward failure. The defect has two components: a saved-output error caused by fused softmax arithmetic and a more fundamental violation of the exact zero-sum score-gradient identity after BF16 casting. The latter creates a mean-key leakage term that grows with key offset and becomes dominant when training produces large keys and nearly one-hot attention.

FA3-SBS repairs the forward channel but cannot repair the cast-induced conservation-law violation. GProj restores the zero row sum using a probability-weighted projection, reducing median μa=∑jpjaj\mu_a=\sum_jp_ja_j0 and μa=∑jpjaj\mu_a=\sum_jp_ja_j1 errors to approximately 0.34% and 0.37%, matching FP32 attention at a reported 4.7% training-step overhead. In the matched 50B-token study, it reaches the same final loss as FP32 attention while avoiding the large-logit regime induced by uncorrected BF16 backward arithmetic. The paper’s broader methodological result is that low-precision kernels should be audited not only against aggregate numerical error, but also against exact structural identities—particularly those that encode invariance and conservation in the underlying derivative.

Paper to Video (Beta)

No one has generated a video about this paper yet.

Whiteboard

Explain it Like I'm 14

1. What is this paper about?

This paper studies a problem in how large AI models are trained.

Modern models often use BF16, a faster but less precise way of storing numbers. They also use a program called FlashAttention to make the attention part of a transformer faster and use less memory.

The researchers found that FlashAttention-3 can behave normally for a long time and then suddenly produce very inaccurate gradients. A gradient is a signal that tells the model how to change its internal numbers to improve. If the gradient is wrong, the model may learn badly.

The surprising part is that the problem can happen without any obvious warning:

  • The training loss gets worse.
  • The gradient becomes extremely large.
  • No NaN values appear.
  • The program does not crash.

The paper explains why this happens and introduces a fix called GProj, short for gauge projection.

2. What questions did the researchers ask?

The researchers wanted to answer several main questions:

  1. Why does FlashAttention-3 become unstable late in training?
  2. Is the problem caused by the model itself, or by inaccurate calculations inside the attention program?
  3. Why can the forward calculation look correct while the backward calculation gives very wrong gradients?
  4. Can the problem be fixed without replacing the fast BF16 system with much slower FP32 calculations?
  5. Does the proposed fix make training both accurate and efficient?

Here, the forward pass means producing the model’s answer. The backward pass works backward from the error and calculates how each internal number should change.

3. How did the researchers investigate the problem?

Training a transformer

They trained a transformer with about 450 million parameters on 50 billion tokens. They compared several versions of attention, including:

  • The normal FlashAttention-3 implementation.
  • A version using more careful calculations in the forward pass.
  • Full or partly higher-precision FP32 attention.
  • Their new GProj method.
  • A method called key smoothing.

They watched the model’s loss, gradient sizes, attention scores, speed, and memory use.

Comparing different levels of numerical precision

The researchers compared BF16 results with results calculated using FP64, a much more precise number format. FP64 was used as a trusted reference, similar to checking a ruler with a very accurate measuring tool.

They measured the difference using relative L2L_2 error. This is a way of asking:

How large is the mistake compared with the correct answer?

For example, an error of 100% means the mistake is about as large as the true gradient itself.

Recomputing only the backward pass

To locate the problem, they kept the forward results exactly the same but recalculated the backward pass in FP32.

This was like keeping the same answer on a test but checking whether the marking instructions were faulty. When they recalculated the backward pass more accurately, most of the extra gradient disappeared. This showed that the main problem was inside the attention backward calculation.

Studying a conservation rule

The researchers examined a mathematical property of softmax attention. In each row, the exact gradient of the attention scores should add up to zero.

This is a kind of conservation law. It is similar to saying that if money is moved between several boxes, the total amount added and removed should balance out.

The exact rule is:

∑jgj=0\sum_j g_j = 0

The paper calls this a zero-sum property.

When the gradient is rounded into BF16, the numbers may no longer add up to exactly zero. Even a very small leftover amount can cause a large error when it is multiplied by large key vectors.

Testing a small example

The researchers also made simple artificial examples with only a few keys. These examples showed that the problem could occur even when the forward pass was perfectly accurate. This helped prove that fixing only the forward pass would not be enough.

4. What did they find?

Finding 1: FlashAttention-3 can fail silently

Training with normal FlashAttention-3 was stable for roughly the first 25 billion tokens. After that:

  • The gradient norm became about 1,000 times larger.
  • Some attention scores grew enormously.
  • The final loss was about 0.2 nats worse than the FP32 version.
  • No NaN values appeared.

This is important because researchers might normally look for crashes or NaNs when debugging training. This problem gives no such clear warning.

Finding 2: There are two separate numerical problems

The paper found two sources of error.

Problem A: A forward-pass rounding mistake

FlashAttention-3 used a fused operation that combined two steps:

  1. Scaling the attention scores.
  2. Subtracting the largest score in the row.

Because these steps were combined and rounded, the largest score was not always changed to exactly zero. This caused the saved attention output to be slightly wrong.

The researchers fixed this by subtracting the maximum score first and scaling afterward. They called this version FA3-SBS, meaning subtract before scale.

This stopped the huge immediate gradient explosion, but it did not completely fix the gradients.

Problem B: The backward pass breaks a zero-sum rule

The deeper problem happens when the score gradient is converted to BF16.

Before rounding, the gradient values in a row add up to zero. After rounding, they might add up to a small nonzero number. The paper calls this leftover amount the row-mass error.

That small leftover is then multiplied by the keys. If the keys have become very large during training, the small error can become a large false gradient.

An analogy is weighing objects on a slightly inaccurate scale. A tiny measurement error may not matter for small objects. But if the error is multiplied by a huge object, the final mistake can become very large.

This problem is especially bad when attention becomes almost one-hot, meaning that one key receives nearly all the attention. In that situation, the true gradient is often very small, so even a small numerical error can become larger than the real answer.

Finding 3: The forward fix alone is not enough

The forward repair, FA3-SBS, made training look more stable, but its query gradients were still highly inaccurate.

The median query-gradient error was about 219%, meaning the error was more than twice the size of the correct gradient. The key-gradient error was about 13%.

This showed that accurate-looking forward outputs do not guarantee accurate gradients.

Finding 4: GProj restores the zero-sum property

GProj fixes the problem by checking how much the rounded gradient fails to add up to zero. It then subtracts a carefully chosen correction so that the gradient once again has a zero sum.

In simplified form, GProj changes the rounded gradient tt into:

t⊥=t−ρmrt^\perp = t - \frac{\rho}{m}r

Here:

  • ρ\rho is the unwanted leftover row sum.
  • rr is a set of BF16 attention probabilities.
  • mm is the sum of those probabilities.
  • t⊥t^\perp is the corrected gradient.

This correction removes the artificial signal caused by the rounding error.

Finding 5: GProj greatly improves gradient accuracy

The table below shows the main comparison:

Method Query-gradient error Key-gradient error
FA3-SBS 219% 13.3%
GProj 0.342% 0.371%
Fused FP32 attention 0.385% 0.371%

GProj therefore produced gradients about as accurate as FP32 attention, even though it continued to use much of the faster BF16 computation.

Finding 6: GProj keeps training stable

In matched training experiments:

  • GProj trained successfully through the full 50 billion tokens.
  • Its final loss matched the FP32 attention version.
  • Normal FlashAttention-3 ended with a worse loss.
  • Key smoothing only delayed the problem; it did not fully solve it.
  • The forward-only repair was stable but still produced inaccurate gradients and very large attention scores.

Finding 7: The fix is relatively cheap

GProj increased the training-step time by about 4.7%.

By comparison, using fused FP32 attention increased the time by about 33.4% in the reported experiment.

So GProj offered a useful compromise:

  • Much better gradient accuracy than ordinary BF16 FlashAttention.
  • Much lower cost than fully using FP32 attention.

5. Why are these results important?

The paper shows that a fast AI calculation can be wrong in a way that is difficult to notice. The model may continue running normally while quietly learning from bad information.

The main lesson is:

It is not enough to check whether the forward answers look correct. The backward gradients must also obey important mathematical rules.

The paper also shows that small numerical errors can become serious when:

  • The model has been training for a long time.
  • Attention becomes very concentrated.
  • The key vectors become large.
  • The true gradient becomes close to zero.

These conditions explain why the problem appears late in training rather than at the beginning.

6. Possible impact of the research

GProj could make large-scale transformer training more reliable while keeping most of the speed and memory advantages of BF16 FlashAttention.

The idea may also be useful beyond this exact attention kernel. Other low-precision computer programs may have similar mathematical rules that are accidentally broken by rounding. Checking and restoring those rules could help prevent hidden training failures.

However, the researchers tested GProj mainly on:

  • BF16 FlashAttention-3.
  • Hopper GPUs.
  • A 450-million-parameter transformer.
  • A particular attention size and sequence length.

More testing is needed to know whether the same problem and solution work equally well for other GPUs, models, number formats such as FP8, and future versions of FlashAttention.

Overall, the paper’s message is simple: fast, low-precision calculations can silently damage learning, but carefully restoring the right mathematical structure can make them both accurate and efficient.

Knowledge Gaps

Knowledge gaps, limitations, and open questions

  • Limited hardware validation: The study evaluates GProj only on Hopper GPUs, particularly H200 hardware; its numerical behavior and performance on Ampere, Ada, Blackwell, AMD, and other accelerators remain unknown.
  • Restricted precision scope: The analysis focuses on BF16 score-gradient casts. It does not establish whether the same conservation-law failure occurs, or how GProj should be adapted, for FP16, FP8, FP4, INT8, stochastic rounding, or mixed-precision formats.
  • Narrow model scale: Training experiments use a single 450M-parameter transformer. It remains unresolved whether the instability, late-training timing, and GProj benefits persist in billion- or trillion-parameter models.
  • Limited architectural diversity: The experiments do not test encoder-only models, encoder–decoder models, vision transformers, multimodal transformers, mixture-of-experts models, alternative positional encodings, or architectures with substantially different attention parameterizations.
  • Restricted attention configuration: The implementation is evaluated with head dimension 64, causal attention, dense and packed layouts, and GQA. Its correctness and overhead for larger head dimensions, multi-head attention without GQA, cross-attention, noncausal attention, sliding-window attention, block-sparse attention, and irregular masks remain untested.
  • Unresolved sequence-length scaling: The reported 4.7% end-to-end overhead is measured at 4096 tokens. The paper does not quantify accuracy, memory use, and runtime as sequence length increases, especially when the additional key-gradient pass becomes dominant.
  • Incomplete value-gradient analysis: GProj corrects query and key gradients but leaves dVdV unchanged. The magnitude, structure, and training impact of value-gradient errors under the same extreme-logit regimes are not characterized.
  • No systematic comparison of projection choices: The method projects along the BF16 probability operand rr, but the paper does not comprehensively compare this choice with projections along exact probabilities, FP32 probabilities, uniform vectors, alternative weighted subspaces, or dynamically selected correction directions.
  • Residual numerical errors remain unexplained: GProj reduces median dQdQ and dKdK errors to approximately the BF16 floor, but several per-capture and synthetic-suite maximum errors remain substantially larger. The sources and practical consequences of these residual outliers are not fully isolated.
  • Inconsistent query/key correction arithmetic: The implemented query and key corrections may use different reconstructed probability operands and separately computed projection coefficients. The effect of this inconsistency on gradient bias, translation invariance, and long-term training has not been quantified.
  • Unclear behavior for zero or very small probability mass: GProj uses a zero-mass fallback, but the numerical stability of ρ/m\rho/m when mm is extremely small, underflowed, or heavily affected by masking is not systematically analyzed.
  • Interaction with masking is incomplete: Although some packed, causal, singleton, and tile-boundary cases are included in the synthetic suite, the paper does not establish guarantees for arbitrary masks, highly sparse supports, variable-length batches, or cross-attention masks.
  • Dependence on saved-output accuracy: The study identifies BF16 saved-output errors as a separate failure channel and repairs the forward softmax with subtract-before-scale. It remains unclear whether GProj alone is sufficient when other saved-state errors, LSE errors, or output-reduction errors are present.
  • No exhaustive decomposition of all backward error sources: The paper separates saved-output error and score-gradient-cast error, but does not provide a complete error budget covering exponentiation, probability reconstruction, reduction order, atomic accumulation, matrix-product rounding, scaling, and final gradient casts.
  • Training conclusions rely on limited replicates: The matched from-scratch results appear to use individual runs for each attention variant. The robustness of final loss, instability onset, and gradient trajectories across random seeds, data orders, and optimizer-state initializations is not established.
  • Limited data and optimization regimes: The experiments use one 50B-token pretraining setup. The effects of learning-rate schedules, batch sizes, optimizers, weight decay, gradient clipping, initialization schemes, token distributions, and curriculum choices remain unexplored.
  • Causal relationship between gradient error and logit growth needs broader validation: The paper provides strong evidence in two layers of one model, but does not determine whether the same feedback loop explains large-logit growth across different layers, models, datasets, or optimization settings.
  • No assessment of downstream task impact: The evaluation focuses on pretraining loss, gradient norms, and numerical accuracy. It does not test whether GProj improves or preserves downstream language modeling, generation quality, calibration, transfer performance, or instruction-tuning outcomes.
  • Long-term convergence is unresolved: Training is reported through 50B tokens, but it remains unknown whether GProj preserves its advantage over substantially longer schedules or whether other numerical instabilities eventually emerge.
  • Comparison with stabilization methods is incomplete: GProj is compared with key smoothing and selected forward repairs, but not systematically with query–key normalization, entropy regularization, activation clipping, optimizer rescaling, stochastic rounding, higher-precision selective recomputation, or combinations of these methods.
  • Generality beyond attention is not demonstrated: The conclusion suggests that conservation-aware auditing may repair other low-precision kernels, but no non-attention example, general algorithm, or empirical validation is provided.
  • Formal guarantees do not cover the implemented kernel fully: The strongest theorem assumes perfect forward quantities and idealized rounding, whereas the implementation uses reconstructed probabilities, BF16 saved outputs, finite-precision reductions, approximate divisions, and separate correction passes. A rigorous bound for the complete implementation is still missing.
  • Adversarial and worst-case inputs remain underexplored: The synthetic suite contains 58 cases and several analytic witnesses, but it does not characterize the full worst-case error over sequence lengths, key offsets, attention sharpness, value ranges, masks, or BF16 rounding patterns.
  • Performance portability is uncertain: The reported timing uses batch one, 4096 tokens, a single H200, and no inter-rank communication. The overhead under realistic large-batch distributed pretraining, tensor parallelism, pipeline parallelism, communication overlap, and different kernel fusion strategies is unresolved.
  • Memory and workspace scaling are not fully characterized: GProj adds query-shaped workspaces and a second key-correction pass, but its peak memory and allocator behavior across batch size, sequence length, head count, and distributed layouts are not reported.
  • Kernel implementation maturity is unclear: The paper does not establish whether GProj has been integrated into an upstream FlashAttention release, whether it remains correct across compiler versions and kernel tiling choices, or whether its numerical behavior is stable under autotuning.
  • Effect of alternative accumulation precision is unknown: The study uses BF16 operands with FP32 accumulators. It does not determine how much of the problem persists with FP64 accumulation, BF16 or FP16 accumulation, TensorFloat-32, FP8 accumulation, or higher-precision selective reductions.
  • The role of stochastic rounding is unresolved: Since the failure is caused by a nonzero row mass after casting, stochastic rounding could alter both the expected bias and variance of the leak. Its interaction with GProj and long-run training stability is not evaluated.
  • Exact conservation is not the only possible invariant: The paper focuses on the zero row sum of the softmax score gradient. Other invariants or symmetries—such as value-translation structure, normalization identities, or invariants induced by grouped-query attention—may also be violated, but are not systematically investigated.
  • Applicability to quantized and compressed attention methods is unknown: The relationship between GProj and FP8/FP4 attention, quantized probabilities, block scaling, per-tensor scaling, and other quantization schemes remains unresolved.
  • No ablation of individual GProj arithmetic components: The separate effects of using the actual BF16 probability mass, FP32 versus lower-precision row sums, the correction placement, the second key pass, and the final correction order are not comprehensively ablated.
  • Practical acceptance criteria are unspecified: The paper reports relative L2L_2 errors and training loss but does not propose deployment thresholds or diagnostic tests for deciding when a low-precision attention kernel is numerically safe in production training.

Practical Applications

Immediate Applications

  • Replace or patch BF16 FlashAttention backward kernels in model training. Implement the paper’s GProj correction in FlashAttention-style kernels, particularly for BF16 attention on Hopper GPUs. The practical workflow is to:
    1. retain the subtract-before-scale forward softmax repair;
    2. measure the row mass of the BF16-cast score gradient, ρ=∑jtj\rho=\sum_j t_j;
    3. measure the mass of the BF16 probability operand, m=∑jrjm=\sum_j r_j;
    4. apply the correction t⊥=t−(ρ/m)rt^\perp=t-(\rho/m)r before the query and key contractions. This is immediately relevant to GPU software, deep-learning frameworks, and foundation-model pretraining. The reported implementation reduced median dQdQ and dKdK errors to approximately FP32 levels with a 4.7% training-step overhead.

Dependencies: Correct integration with the target kernel’s tiling, masking, grouped-query attention, BF16 casting, and FP32 accumulation behavior. The reported cost and accuracy are demonstrated primarily for Hopper GPUs, head dimension 64, causal attention, and a 450M-parameter transformer.

  • Add numerical-fidelity tests to attention-kernel continuous integration.
    • common translations of all keys;
    • nearly one-hot attention distributions;
    • large key offsets with small key-to-key differences;
    • unequal and singleton attention supports;
    • packed and tile-boundary layouts;
    • long-sequence causal attention.

These tests can detect silent gradient corruption even when loss, activations, and outputs appear numerically normal.

Dependencies: A trustworthy high-precision reference implementation, representative adversarial inputs, and tolerances that distinguish ordinary BF16 rounding from catastrophic conservation-law violations.

  • Use gradient-conservation diagnostics during large-scale pretraining.
    • attention logit scale;
    • maximum attention probability;
    • query/key norms;
    • gradient norm by layer;
    • deviation from key-translation invariance.

A rising ρ\rho in layers with large key norms and highly concentrated attention can serve as an early-warning signal for silent late-training failure. This is applicable to training infrastructure, observability platforms, and distributed experiment management.

Dependencies: Access to kernel-level intermediate statistics without causing excessive synchronization or memory traffic. Monitoring alone does not repair the gradient.

  • Adopt subtract-before-scale softmax computation as a separate forward repair. The paper identifies a fused multiply-add issue in the forward softmax and recommends subtracting the unscaled row maximum before applying the attention scale. This can be deployed independently of GProj to reduce saved-output errors and avoid extreme-input failures.

Dependencies: This repair addresses the saved-output channel but not the post-cast score-gradient leak. It should therefore not be treated as a substitute for backward projection.

  • Audit existing training runs for silent attention-kernel failures.
    • native BF16 attention;
    • FP32 attention;
    • GProj or another conservation-aware backward;
    • high-precision reference gradients on selected layers.

If recomputing only a few attention backward passes substantially lowers the gradient norm, the training run may have been affected by numerical kernel error rather than by data, optimizer, or model-design problems.

Dependencies: Saved inputs, model checkpoints, reproducible kernels, and sufficient compute to replay representative batches. The paper’s diagnosis relied on layer-specific intervention while holding forward outputs fixed.

  • Improve debugging workflows for unexplained late-training instability.
    • attention becoming nearly one-hot;
    • unusually large query/key norms;
    • common key offsets;
    • disagreement between BF16 and FP32 attention backward passes.

This provides a concrete alternative to immediately changing learning rates, clipping thresholds, normalization, or optimizer settings.

Dependencies: The symptoms may also arise from genuine optimization or data problems. The diagnostic should be used alongside, not instead of, standard training checks.

  • Use the conservation-law principle in academic benchmarking of low-precision kernels. Researchers evaluating BF16, FP8, FP4, or other fused attention implementations can test whether exact algebraic identities survive quantization. In attention, the key identity is the zero row sum of the softmax score gradient. More generally, researchers can search for invariants such as normalization, orthogonality, conservation, or symmetry constraints and evaluate whether quantized operands preserve them.

Dependencies: The relevant invariant must be derived for the specific operation and must be checked on the operands actually consumed by the low-precision matrix multiplications, not only on higher-precision logical quantities.

  • Use GProj as a lower-cost alternative to fully FP32 attention in selected production training workloads. Teams that currently switch to fused FP32 attention to recover training stability may use GProj to retain BF16 operands and most of the FlashAttention performance profile. In the reported experiment, fused FP32 attention added approximately 33.4% to the training-step time, whereas GProj added approximately 4.7%.

Dependencies: The performance advantage may change with sequence length, head dimension, GPU architecture, attention fraction of the total step, and implementation quality. The paper explicitly reports limited hardware and model coverage.

  • Incorporate translation-invariance checks into compiler and kernel validation. Compiler, GPU-library, and accelerator teams can construct paired inputs in which every key is shifted by the same vector. The exact query gradient should remain unchanged, while an uncorrected low-precision implementation may change in proportion to the residual row mass. This creates a simple black-box regression test for fused attention kernels.

Dependencies: The test must hold the relevant incoming derivative fixed and distinguish changes caused by legitimate representation differences from changes caused by gradient leakage.

Long-Term Applications

  • Generalize conservation-aware projections to FP8, FP4, and other quantized attention formats.
    • FP8 attention and training;
    • FP4 or mixed FP4/FP8 attention;
    • quantized recurrent or state-space models;
    • sparse and block-sparse attention;
    • mixture-of-experts routing probabilities.

A future library could expose a general conservation_aware_backward abstraction rather than a BF16-specific GProj implementation.

Dependencies: Each datatype has different rounding, underflow, saturation, and scaling behavior. The projection may need datatype-specific probability operands, scaling rules, or stochastic-rounding analysis.

  • Develop a generic compiler pass for invariant-preserving automatic differentiation. Automatic differentiation systems could annotate operations with algebraic constraints—such as zero-sum, normalization, or gauge invariance—and insert low-cost projections after quantization or casting. For attention, the compiler would automatically preserve the score-gradient zero-sum condition before contraction with keys.

Dependencies: The compiler must infer or receive valid invariants, preserve masking and sparsity semantics, and estimate whether the added reductions and correction passes justify their runtime cost.

  • Design conservation-aware hardware primitives for accelerators.
    • row-mass reductions;
    • normalized rank-one corrections;
    • invariant-preserving matrix products;
    • fused projection and contraction;
    • low-precision reductions with controlled mass error.

Such primitives could reduce the extra pass and workspace currently required for the dKdK correction.

Dependencies: Hardware support would require evidence that the failure occurs broadly across models and precisions, along with careful area, power, scheduling, and memory-bandwidth analyses.

  • Create reliability standards for low-precision training kernels.
    • gradient error against FP64 references;
    • behavior under key/value translations;
    • performance under sharp attention;
    • stability over long training horizons;
    • sensitivity to sequence length and model scale;
    • silent failure rates without NaNs.

This could inform procurement, reproducibility requirements, and release criteria for open-source GPU kernels and commercial AI platforms.

Dependencies: Standards bodies and vendors would need agreed reference implementations, representative workloads, and reporting protocols that do not unfairly penalize valid hardware-specific optimizations.

  • Improve reproducibility and comparability of foundation-model training. Training recipes could report the attention kernel, precision of each forward and backward component, conservation-law tests, and whether gradient projections are enabled. This would make results from different GPU stacks and FlashAttention versions more comparable and could explain otherwise unexplained differences in loss or stability.

Dependencies: Reproducibility also depends on data ordering, optimizer state, random seeds, compiler versions, and distributed reduction behavior. GProj would address only one class of numerical discrepancy.

  • Apply invariant-preserving numerical methods to other machine-learning operations.
    • probability normalization in routing and mixture-of-experts layers;
    • zero-sum gradients in normalized losses;
    • orthogonality constraints in representation learning;
    • conservation laws in differentiable physics;
    • equivariant operations in geometric deep learning;
    • mass-conserving models for fluid, climate, or energy simulation.

A possible product or research workflow is an invariant audit that automatically identifies whether quantization introduces forbidden components into a gradient.

Dependencies: Transfer requires proving that the invariant is exact for the intended computational graph and determining a projection that does not distort legitimate gradient information.

  • Enable safer long-context and highly concentrated-attention training. Since the leak is amplified by large key norms and sharp, nearly one-hot attention, conservation-aware kernels may become increasingly important for long-context transformers, retrieval-augmented models, attention-sink architectures, and models with large activation outliers.

Dependencies: The paper evaluates 4,096-token sequences and a 450M-parameter model. Longer contexts, larger models, different positional encodings, and distributed attention layouts may introduce additional numerical failure modes not addressed by GProj.

  • Support numerical-risk-aware optimizer and training-schedule policies. Training platforms could dynamically switch attention implementations or precision levels when diagnostics detect large logit scales, high attention concentration, or rising score-gradient row mass. For example, a system might use BF16 GProj by default and temporarily fall back to FP32 attention for layers or batches entering an unsafe regime.

Dependencies: Dynamic switching introduces implementation complexity, possible nondeterminism, additional profiling overhead, and the risk of masking rather than eliminating underlying kernel errors.

  • Inform daily-use AI systems indirectly through more reliable model training. Although GProj is not a consumer-facing algorithm, more stable and accurate low-precision training could improve the reliability of deployed language, vision, recommendation, and multimodal systems while reducing training cost and energy consumption. Potential downstream benefits include fewer failed pretraining runs, more consistent model quality across hardware platforms, and reduced need for expensive FP32 fallback training.

Dependencies: These benefits are indirect and depend on successful integration into production kernels, validation across larger models and multiple accelerator generations, and confirmation that the reported stability gains persist at industrial scale.

Glossary

  • Attention sink: A token or position that receives disproportionately high attention, often regardless of content. “the near one-hot pattern also seen in attention sinks and massive activations”
  • BF16 (bfloat16): A 16-bit floating-point format with a wide exponent range and reduced precision. “BF16 is now standard in large-scale pretraining”
  • Bitwise identical: Exactly the same at the level of machine-represented bits. “while keeping every forward activation bitwise identical”
  • Causal support: The set of positions that a query is permitted to attend to under a causal attention mask. “let LL be its causal support”
  • Conservation law: An invariant quantity preserved by a mathematical operation or transformation. “The remaining error comes from a broken conservation law.”
  • Covariance form: A representation of a quantity using deviations from weighted means, making translation invariance explicit. “which gives the covariance form”
  • Dynamic softmax shifting: A numerical technique that adjusts softmax inputs during computation to improve stability. “Related remedies followed, from a dynamic softmax shift”
  • Eager attention: An unfused, explicitly executed implementation of attention operations. “FP32 attention (eager)”
  • Fused multiply-add (FMA): A hardware operation that multiplies two values and adds a third with a single rounding step. “a fused multiply-add in the forward softmax”
  • Fused kernel: A GPU kernel that combines multiple computational operations to reduce memory movement and improve performance. “attention computed by fused kernels such as FlashAttention”
  • Gauge projection: A projection that removes a component associated with a symmetry or redundant coordinate, here enforcing a zero row sum. “We introduce GProj (gauge projection)”
  • Gradient norm: The magnitude of a model’s gradient, commonly used to monitor optimization stability. “the gradient norm grew a thousandfold”
  • Grouped-query attention (GQA): An attention configuration in which multiple query heads share key and value heads. “over the query heads that share a KV head in grouped-query attention (GQA)”
  • Head dimension: The dimensionality of the query, key, and value vectors within an attention head. “with head dimension 64”
  • Hopper GPU: A generation of NVIDIA GPU hardware designed for high-performance computing and machine learning. “The GProj kernel targets Hopper GPUs”
  • Interquartile range: The range between the 25th and 75th percentiles of a distribution. “median and interquartile range over 256 rows”
  • Key smoothing: A stabilization method that subtracts an average or common component from attention keys. “Key smoothing delays the onset by about 3B tokens”
  • Log-sum-exp (LSE): A numerically stable way to compute the logarithm of a sum of exponentials. “the FP32 log-sum-exp (LSE) as usual”
  • Logit: An unnormalized score used as input to a softmax function. “training still drives attention logits to thousands of times their size”
  • Low-precision arithmetic: Computation using numerical formats with fewer bits than standard floating-point representations. “Low-precision arithmetic is what makes large-scale pretraining affordable”
  • Massive activation: An unusually large neural-network activation that can affect numerical stability and optimization. “the near one-hot pattern also seen in attention sinks and massive activations”
  • Matrix product accumulated in FP32: A lower-precision matrix multiplication whose partial sums are accumulated using 32-bit floating point. “a BF16 matrix product accumulated in FP32”
  • Mixed-precision training: Training that combines multiple numerical precisions for efficiency and numerical stability. “Mixed-precision and BF16 training”
  • One-hot pattern: A distribution in which nearly all probability mass is concentrated on a single element. “almost every query puts nearly all of its attention on a single key, the near one-hot pattern”
  • Outlier activation: An activation whose magnitude is much larger than typical activations. “or outlier activations”
  • Packed layout: A memory arrangement that stores multiple variable-length or structured sequences compactly. “dense and packed causal layouts”
  • Pre-clipping gradient norm: The gradient magnitude measured before gradient clipping is applied. “pre-clipping gradient norm”
  • Rank-one correction: An adjustment expressible as the outer product of two vectors, affecting a matrix through a single-dimensional component. “with two rank-one corrections per row”
  • Relative L2L_2 error: The Euclidean error between an approximation and reference, normalized by the reference magnitude. “Full-tensor relative L2L_2 error (\%)”
  • Row mass: The sum of the entries in a row of a score-gradient or probability-related tensor. “The exact score gradient has a conserved quantity: its row mass”
  • Rounding model: A mathematical abstraction describing the error introduced when numerical values are rounded to a finite-precision format. “Over cast errors allowed by the rounding model”
  • Saved output: An intermediate forward-pass result retained for use during backpropagation. “FA3's backward does not recompute μa\mu_a; it estimates it as δ^\widehat\delta”
  • Score gradient: The derivative of the loss with respect to the attention score matrix. “The score derivative gg is one row of dSdS”
  • Softmax: A function that converts a vector of scores into a probability distribution by exponentiating and normalizing its entries. “The softmax score gradient sums to zero along every row”
  • Stochastic rounding: A rounding method that randomly selects neighboring representable values according to their distances from the exact value. “which motivated stochastic rounding and careful BF16 recipes”
  • Tiled pseudocode: Algorithmic notation describing computation over blocks or tiles of a larger tensor. “Algorithm~\ref{alg:gproj-main} gives the tiled pseudocode”
  • Translation invariance: The property that a result remains unchanged when a common offset is added to related inputs. “For fixed t,rt,r, this contraction is unchanged by any common translation of the keys.”
  • Unit roundoff: A bound characterizing the maximum relative error introduced by rounding in a floating-point system. “where ub=2−8u_b=2^{-8} is the BF16 unit roundoff”
  • Vector–Jacobian product (VJP): The product of a vector with the Jacobian, commonly used to compute reverse-mode autodifferentiation gradients. “the vector--Jacobian product (VJP) is”
  • Zero-sum subspace: The set of vectors whose components sum to zero. “The result, t⊥=t−λrt^\perp=t-\lambda r, is the projection of tt onto the zero-sum subspace along rr.”

Open Problems

We haven't generated a list of open problems mentioned in this paper yet.

Tweets

Sign up for free to view the 3 tweets with 478 likes about this paper.