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.
Paper Prompts
Sign up for free to create and run prompts on this paper.
Top Community Prompts
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
NaNvalues 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:
- Why does FlashAttention-3 become unstable late in training?
- Is the problem caused by the model itself, or by inaccurate calculations inside the attention program?
- Why can the forward calculation look correct while the backward calculation gives very wrong gradients?
- Can the problem be fixed without replacing the fast BF16 system with much slower FP32 calculations?
- 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 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:
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
NaNvalues 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:
- Scaling the attention scores.
- 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 into:
Here:
- is the unwanted leftover row sum.
- is a set of BF16 attention probabilities.
- is the sum of those probabilities.
- 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 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 , 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 and 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 when 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 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:
- retain the subtract-before-scale forward softmax repair;
- measure the row mass of the BF16-cast score gradient, ;
- measure the mass of the BF16 probability operand, ;
- apply the correction 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 and 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 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 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 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 error: The Euclidean error between an approximation and reference, normalized by the reference magnitude. “Full-tensor relative 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 ; it estimates it as ”
- Score gradient: The derivative of the loss with respect to the attention score matrix. “The score derivative is one row of ”
- 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 , 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 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, , is the projection of onto the zero-sum subspace along .”



