Conceptio › Archive › arXiv CS
arXiv CSopen access

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

Junlin Chen et al. · arxiv_cs
arXiv CS · Papers · License: Open Access
Open Source ↗Direct PDF ↓
clouddistributed-computingparallel-computing
distributed computing, parallel computing, cloud

Preprint

arXiv:2609.34272v1 [cs.LG] 28 Sep 2026

B ROKEN S YMMETRY IN BF16 ATTENTION : W HY F LASH ATTENTION G RADIENTS B LOW U P L ATE IN T RAINING Junlin Chen1,5 Daize Dong1 Huanwei Di1 Haolong Jia1 Jiawei Wu1 Haotian Xie1 Mingkai Zheng1 Yang Li1 Leshang Chen2 Huishu Wang3 Eric P. Xing4,5 Hongyi Wang1,5 1 Rutgers University 2 Oracle 3 New York University 4 MBZUAI 5 Carnegie Mellon University [email protected]

A BSTRACT 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.

1

I NTRODUCTION

Low-precision arithmetic is what makes large-scale pretraining affordable (Micikevicius et al., 2018; Kalamkar et al., 2019). Frontier and fully open models are pretrained with BF16 or FP8 matrix products over trillions of tokens (Grattafiori et al., 2024; DeepSeek-AI, 2024; Kimi Team et al., 2025; IFM Team, 2026), with attention computed by fused kernels such as FlashAttention (Dao et al., 2022; Dao, 2024; Shah et al., 2024) that keep only BF16 operands and a few FP32 statistics. These kernels are validated against a higher-precision reference on random, well-conditioned inputs and then trusted in training. The trust is rarely tested where it matters, because a wrong gradient does not announce itself: it causes no crash and no NaN, and the worse model it produces is easily blamed on hyperparameters or data. Instabilities that do surface are usually traced to the model or the optimizer, such as growing attention logits (Dehghani et al., 2023; Wortsman et al., 2024), attention entropy collapse (Zhai et al., 2023) or outlier activations (Fishman et al., 2025); the numerical fidelity of the attention kernel itself is rarely examined (Golden et al., 2024). We ran into such a case while pretraining a 450M-parameter transformer on 50B tokens with FlashAttention-3 (FA3). For the first 25B tokens, training was healthy. Then the gradient norm grew a thousandfold, and the loss drifted up and never recovered, ending 0.2 nats above an otherwise identical run with FP32 attention (Section 5). Nothing overflowed and nothing became NaN. 1

Preprint

To find the source, we recomputed the attention backward of the two most affected layers in FP32 while keeping every forward activation bitwise identical; the excess gradient almost disappeared (Section 3). The fault was in the attention backward, which raises the question this paper answers: how can a kernel whose outputs are accurate return gradients that are this wrong? Part of the answer is already known. The community has traced a BF16 NaN in attention to a fused multiply-add in the forward softmax: to save an instruction, “scale” and “subtract the row maximum” are merged into one rounded operation, so the largest score no longer maps exactly to zero (PyTorch contributors, 2024). Because the report involved extreme logits, the fix was reasonably treated as a corner case: it became an option in FlashAttention-2, off by default, and never reached FA3 (FlashAttention contributors, 2024). Related remedies followed, from a dynamic softmax shift (Qiu & Yao, 2026) to keeping the attention output in FP32 (Kimi Team, 2026). When we repair the forward in the same spirit, subtracting the maximum before scaling (FA3-SBS), the thousandfold growth disappears and training stays stable to the end. The known fix seemed to be enough. It was not. With the forward repaired, the query gradients of the affected layers are still wrong by more than their own size (Section 5), and the run still drives these layers into extreme attention: by 33.6B tokens their typical attention score is about 4–5 × 104 , against under 10 in the same model trained with accurate gradients (Section 3.1). Stable training, it turns out, is not evidence of correct gradients, and the remaining error has a different origin, one that no forward repair can reach. That origin is a broken symmetry. The gradient with respect to a query should depend only on how the keys differ from one another, never on where they sit as a group. Every exact attention backward enforces this translation symmetry through a conservation law: the softmax score gradient sums to zero along each row, and that zero sum is what cancels the common part of the keys. FA3 rounds this score gradient to BF16 before multiplying it with the keys. The rounding leaves a small nonzero row sum, and the product turns it into a spurious term proportional to the mean key, which we call the leak (Figure 1, panel 1). The leak grows with how far the keys sit from the origin relative to their spread (panel 2), and when attention also concentrates on a single key, the true gradient is nearly zero, so the leak outweighs the signal by orders of magnitude. Because the error is created after the forward pass has finished, a four-key example with an exact forward already exhibits it (Section 3.4). This view also explains why the existing remedies stop halfway. The FMA fix, FP32 outputs and dynamic shifting all act on the forward pass. Key smoothing (Zhang et al., 2026) subtracts the average key before the kernel, which shrinks the mean key the leak is multiplied by but leaves the broken sum intact. When attention concentrates on a single key, the mean that matters is that key itself, which smoothing does not remove: in our runs, key smoothing delays the blow-up by about 3B tokens but does not prevent it. The natural fix is therefore to restore the law itself. GProj (gauge projection) puts the zero row sum back after the cast, using the actual mass of the BF16 probabilities the kernel multiplies, with two (1)(1) The castcast leaks alongΣμ error (%) BF16 breaks (2) GProj projects it bac j gkj = 0 (2) Query-gradient exact g: Σ = 0 keys query gradient BF16 t: Σ=ρ≠0 spread σ (small) k

1

102

FA3-SBS

ρμ

k

(la rge

lea kα

−1 0

0

10 error αρ μk = −24 (exact dq = 0) 10−1

μk

0

)

101

k1

k2 exact k3≈ GProj k4

FA3 forward repair

dQ =GProj α tK leaks αρ μk

(3) Training gradient norm FA3 FA3-SBS GProj

103 102 101

GProj

dQ = α (t − λr)K

100

λ = ρ/m: no leak

10−1

r: BF16 probabilities, m=3 101 100 102 10

key offset ‖μk‖/σk

0

25B

training tokens

50B

Figure 1: (1) The exact score gradient sums to zero along each row; its BF16 cast t leaves a row mass ρ ̸= 0, which adds a leak αρµk along the mean key to the query gradient, while GProj removes ρ after the cast. (2) Synthetic rows with an exact forward and only the score gradient cast to BF16: the error with a perfect forward grows with the key offset, so no forward repair can remove it, while GProj stays at the BF16 floor (median and interquartile range over 256 rows; Appendix C). (3) FA3 blows up late in training; GProj does not (Section 5). 2

Preprint

rank-one corrections per row: one inside the existing backward pass and one in a second pass for the key gradient. It brings the query and key gradient errors down to the level of FP32 attention, for 4.7% more time per training step. Trained from scratch, it stays stable throughout and ends at the same loss as FP32 attention. In summary, this paper makes three contributions: • We identify a silent late-training failure of BF16 FlashAttention-3 and trace it to the attention backward. • We explain it through a conservation law of the attention gradient that the BF16 cast breaks, an error no forward repair can remove. • We propose GProj, which restores the law and matches FP32 attention in gradient accuracy and final loss at 4.7% extra cost.

2

R ELATED W ORK

Low-precision training and attention kernels. Mixed-precision and BF16 training (Micikevicius et al., 2018; Kalamkar et al., 2019) and, more recently, FP8 training (Peng et al., 2023; Fishman et al., 2025) trade precision for throughput; rounding can bias updates, which motivated stochastic rounding and careful BF16 recipes (Gupta et al., 2015; Zamirai et al., 2020). FlashAttention and its successors reduce attention memory traffic (Dao et al., 2022; Dao, 2024; Shah et al., 2024), and quantized attention pushes precision further down (Zhang et al., 2025a;b; Cheng et al., 2025). Golden et al. (2024) quantify the numerical deviation FlashAttention introduces during training, mainly through its forward outputs; we find a backward error that can dominate the gradient. The FMA cancellation hazard behind FA3-SBS is known (PyTorch contributors, 2024); FA3-SBS repairs the saved output, not the score-gradient cast. Softmax reformulations. Beyond classical shifted-softmax analysis (Blanchard et al., 2021; Pébay, 2008), Qiu & Yao (2026) connect saved-output error to the backward reduction and biased updates, motivating conditional dynamic shifting. Kimi K3 adopts the related remedy of keeping the attention output in FP32 during training (Kimi Team, 2026, Section 2.1). Both address only the saved-output channel. Conservation-aware low-precision kernels. SageBwd uses the exact zero row sum of the score gradient to motivate smoothing keys before quantization (Zhang et al., 2026, Sections 4.2 and 6). Direct-P matches normalization to consumed probabilities for FP4 (Hu, 2026, Section 4.3), and MXAttention normalizes by the mass of the quantized exponentials (Yu et al., 2026). GProj instead measures the mass after the BF16 cast and projects the contractions with ρ/m (Section 4). Training-stability interventions. Growing attention logits are a known source of instability, addressed by query–key normalization (Henry et al., 2020; Dehghani et al., 2023; Wortsman et al., 2024), entropy control (Zhai et al., 2023), or direct rescaling of query and key weights during optimization (Kimi Team et al., 2025; Liu et al., 2025; Heo et al., 2021). These methods change the model or the optimizer; we show that the logit growth itself can be driven by the kernel’s gradient error, which GProj repairs.

3

W HERE THE G RADIENT E RROR C OMES F ROM

A wrong gradient is hard to see, because nothing in the forward pass changes. In the FA3 run of Figure 3, a 450M-parameter transformer (Appendix H), loss and gradient norm look normal for 25B tokens; the gradient norm then passes 10 by 30B tokens and the loss degrades, with no NaN at any point (Figure 2). This section first shows that the excess gradient is a numerical error of the attention backward, and then derives where that error comes from. 3

Preprint

3.1

T HE EXCESS GRADIENT IS A NUMERICAL ERROR

The failure is confined to√two layers. We measure the typical score magnitude of a layer by its logit scale α rms(Q) rms(K) d, with α the attention scale and d the head dimension. It stays moderate in 22 of the 24 layers but grows by orders of magnitude in layers 5 and 11, where queries and keys become large and almost every query puts nearly all of its attention on a single key, the near one-hot pattern also seen in attention sinks and massive activations (Xiao et al., 2024; Sun et al., 2024; Gu et al., 2025). Large logits alone do not make a gradient wrong, so we separate the two directly. We freeze the model at 33.6B tokens and recompute only the attention backward of these two layers in FP32, leaving every attention output and the loss bitwise identical. The model-gradient norm falls by more than two orders of magnitude (Figure 2, panel 3): almost all of the observed gradient is numerical error from two attention backward passes. The error and the large-logit regime feed each other. In the matched from-scratch runs (Section 5), which change only the attention computation, every variant whose backward keeps the BF16 error, including the forward repair of Section 3.3 and key smoothing, drives layers 5 and 11 into this regime, while FP32 attention and GProj stay out of it (Table 3). Since FA3-SBS and GProj differ only in the backward, the gradient error itself drives the logits up; larger logits in turn come with larger keys, which amplify the error (Theorem 4.2). 3.2

A CONSERVATION LAW OF THE ATTENTION BACKWARD

For one query q ∈ Rd , let L be its causal support, kj ∈ Rd and vj ∈ Rdv the keys and values, α > 0 the attention scale, and u = ∂L/∂o the incoming derivative. Define sj = αq ⊤ kj , aj = u⊤ vj ,

µa =

X

pj = P pj aj ,

esj ℓ∈L e

µk =

j

, s ℓ

o=

X

pj vj ,

j

X

pj kj ,

(1)

gj = pj (aj − µa ).

j

Sums run over L. The score derivative g is one row of dS, and the vector–Jacobian product (VJP) is X GQ = dq = α gj kj , GK,j = dkj = αgj q, GV,j = dvj = pj u. (2) j

Capital Q, K, V, U collect rows, and dQ, dK, dV are full gradient tensors; key and value contributions also sum over queries and over the query heads that share a KV head in grouped-query attention (GQA). We measure errors against this VJP in FP64 on the same BF16 inputs. The exact score gradient has a conserved quantity: its row mass, the sum of its entries, is zero. This is what makes attention gradients insensitive to where the keys sit as a group. Proposition 3.1 (Exact translation structure). At fixed q,P u, common translations of unmasked keys or values preserve the local VJP; 1⊤ g = 0 and GQ = α j pj (kj − µk )(aj − µa ). P P P The proof is one line: j gj = j pj aj − µa j pj = 0, so shifting every key by µk leaves P g k unchanged, which gives the covariance form (Appendix A). A numerical backward can j j j break this law in two places: before the score gradient is formed, through the quantities it is built from, or after, when it is rounded and multiplied with the keys. We call these channels one and two and examine them in turn. 3.3

C HANNEL ONE : THE SAVED OUTPUT

FA3’s backward does not recompute µa = u⊤ o; it estimates it as δb = u⊤ ob (the row statistic D in Algorithm 1) from the BF16 output ob saved by the forward pass. An error eδ = δb − µa in this single reduction shifts the score gradient by −eδ p and produces the query-gradient error Eδ = −αeδ µk , 4

(3)

2.1 2.0 1.9 1.8 25

30

Training tokens (B)

(2) Gradient norm rises 103 102 101 10

0

25

30

Full-model gradient norm

(1) Loss degrades

Pre-clipping gradient norm

Training loss (10-update mean)

Preprint

(3) Same forward, different backward in layers 5 and 11 104

103

102

101 FA3

Training tokens (B)

Correct O, Correct O, FP32 row 0 all rows backward

Figure 2: The failure and its localization. (1)–(2) Loss and pre-clipping gradient norm of the FA3 run of Figure 3 (10-update means). (3) Model-gradient norm on eight documents at 33.6B tokens, with the forward held bitwise fixed and only the backward of layers 5 and 11 changed: FA3’s own; FA3 given a correct saved output (“correct O”, recomputed in FP64 and rounded to BF16) for the first query row or for all rows (removes channel one, Section 3.3); or FP32 (removes both channels). a vector along the attention-weighted mean key. This one-line prediction matches the measured error almost exactly, and supplying a correct saved output removes most of it, together with nearly all of the excess model gradient (Figure 2; Appendix I). The saved-output error has a known source: a fused multiply-add in the forward softmax merges scaling and maximum subtraction into one rounded operation, so the largest score no longer maps exactly to zero (PyTorch contributors, 2024). Subtracting the unscaled row maximum before scaling, a one-line change we call FA3-SBS (subtract before scale), removes it and substantially lowers the gradient errors (Section 5). 3.4

C HANNEL TWO : THE CAST BREAKS THE CONSERVATION LAW

Removing channel one repairs the model-gradient norm but not the query gradient. With a correct saved output the norm returns to its FP32 level, yet the query gradient is still wrong by more than its own size and poorly aligned with the exact one (Table 10), and this residual still drives layers 5 and 11 into the large-logit regime. It enters after the score gradient is formed: FA3 rounds g to a BF16 operand t beforeP contracting it with the keys, and rounding need not preserve a zero sum. With e = t − g and ρ = j tj , the contraction satisfies, before its own rounding, X X α tj kj − GQ = α ej (kj − µk ) + αρµk . (4) j

j

The first term is an ordinary rounding error, measured relative to the mean key. The second is a leak: the row mass times the mean key itself. It is harmless while keys are small and attention is spread out, and it dominates late in training. When a row puts almost all of its weight on one key, the true gradient is a covariance over a nearly one-hot distribution and is close to zero, whereas ∥µk ∥ is close to the norm of that key, which late training makes large; a row mass of a single BF16 rounding step then outweighs the signal by orders of magnitude. Two keys make this explicit. With weights 1 − ε and ε, GQ = α ε(1 − ε)(a1 − a2 )(k1 − k2 ), |ρ| ≤ 2ub ε(1 − ε)|a1 − a2 |, (5) where ub = 2−8 is the BF16 unit roundoff, so the rounding model allows a leak of up to 2ub ∥µk ∥/∥k1 − k2 ∥ relative to the true gradient. In this bound the sharpness ε cancels: concentrating attention does not protect the gradient, while every increase in key norm relative to key separation enlarges the leak linearly. The four-key example below realizes a leak of this kind under actual rounding. Neither FP32 accumulation nor a forward repair can remove the leak: it isP already in the BF16 multiplicands, and it arises after the forward pass. The key gradient dkj = α i tij qi contracts the same operand along the queries and inherits its row-mass error. A four-key example. Take q = 0, so attention is uniform over four keys kj = (65536, j − 1) that share their first coordinate; the first coordinate of the true query gradient is therefore exactly zero. With α = 1/8 and values chosen so that g = (1 + η, −1, −η, 0), where η = 3/1024, every input and the saved output are exact in BF16. Casting g to BF16 gives t = (1, −1, −η, 0): the entry 5

Preprint

Algorithm 1: FA3 backward pass (black) with the GProj additions (orange) Input: BF16 Q, K, V, O, dO ∈ RN ×d ; FP32 LSE ∈ RN ; scale α; key tiles Kj , Vj , query tiles Qi , Oi , dOi . Output: BF16 dQ, dK, dV . Accumulators and row statistics are FP32. 1 D ← rowsum(dO ⊙ O); dQacc ← 0; B acc ← 0, ρ ← 0, m ← 0 2 for each key tile j in parallel do 3 Load Kj , Vj ; dKjacc ← 0, dVj ← 0 4 for each query tile i that attends to tile j do 5 Load Qi , dOi , LSEi , Di ; S ← α Qi Kj⊤ ; P ← exp(S − LSEi ) (masked) ⊤ 6 dVj += BF16(P )⊤ dOi ; dP  ← dOi Vj 7 t ← BF16 P ⊙ (dP − Di ) 8 r ← BF16(P ); ρi += rowsum(t); mi += rowsum(r) 9 dQacc += t Kj ; Biacc += r Kj (atomic adds) i 10 dKjacc += t⊤ Qi 11 end for 12 Write dVj ; FA3 writes dKj ← BF16(α dKjacc ), GProj keeps the FP32 dKjacc for Pass 2 13 end for  14 λ ← ρ/m (zero where m = 0); dQ ← BF16 α(dQacc − λ ⊙ B acc ) 15 Pass 2: for each key tile j in parallel do C ← 0 16 for each query tile i that attends to tile j do ⊤ 17 Recompute P as in line 5; re ← BF16(P  ); C += re BF16(λi ⊙ Qi ) acc 18 end for; dKj ← BF16 α(dKj − C)

1 + η rounds to 1, leaving row mass ρ = −η. The first coordinate of the query gradient becomes αρ · 65536 = −24 instead of 0 (Appendix A lists the values).

4

GP ROJ

To remove the row-mass leak of the BF16 score-gradient cast (Section 3.4), GProj restores the zero row sum after the cast, on the operands the kernel actually multiplies. It removes the row mass ρ that the cast leaves in t by subtracting a multiple ofPthe BF16 probabilities r, an operand the kernel already forms for dV , whose actual mass is m = j rj . Subtracting λr leaves row mass ρ − λm, so cancellation requires λ = ρ/m; using ρ alone leaves ρ(1 − m), because BF16 probabilities need not sum to one. The result, t⊥ = t − λr, is the projection of t onto the zero-sum subspace along r. It never has to be formed, because both contractions are linear in it:  ⊤ dK = α t⊥ Q = α t⊤ Q − r⊤ (λ ⊙ Q) .

 dQ = α t⊥ K = α tK − λ ⊙ (rK) ,

FA3 already computes tK and t⊤ Q, so GProj adds, per gradient, one product (rK, and r⊤ (λ ⊙ Q) in a second pass) and one rank-one term per query row, plus two row sums, keeping BF16 products, FP32 accumulators, and the existing saved BF16 output and FP32 LSE. The query correction fits in FA3’s own backward pass. The key correction needs a second pass: λ is defined per query row and is known only after that row has seen every key, whereas dK is accumulated per key tile over all query rows. GProj keeps the subtract-before-scale forward of FA3-SBS. Algorithm 1 shows where GProj enters FA3’s backward, and kernel details are in Appendix B. The projection works because the exact score gradient already has zero row sum: since 1⊤ g = 0, the projection leaves g unchanged and acts only on the cast error e = t − g, removing its row mass. In the query contraction this subtracts the r-weighted center µr from every key, and a common key offset cancels exactly. Proposition 4.1 states this precisely; proofs for this section are in Appendix C. Proposition 4.1 (Projection with consistent operand mass). For nonnegative r with m > 0, define Πr = I − r1⊤ /m and λ = ρ/m. Then Π2r = Πr , 1⊤ Πr = 0, and t⊥ = Πr t = t − λr has zero row 6

Preprint

sum. With µr =

P

j rj kj /m, its query contraction is

 X

Gproj = α Q

j

 X ρ X tj kj − rj kj  = α tj (kj − µr ). m j j

(6)

For fixed t, r, this contraction is unchanged by any common translation of the keys. If e = t − g for the exact score gradient g, then X t⊥ − g = Πr e, Gproj − GQ = α ej (kj − µr ). (7) Q j

It is also the smallest zero-sum correction in the r-weighted norm: t−λr minimizes over all z with 1⊤ z = 0.

2 j (zj −tj ) /rj

P

The remaining error is therefore bounded by the spread of the keys, not by their offset; a backward that contracts t directly has no such bound. To compare the two on equal terms, give every method a perfect forward: exact probabilities p and the exact center µa , so the pre-cast score operand is exactly x = g. Only the P BF16 cast remains, tj = xj + Pξj with |ξj | ≤ ub |xj | for the unit roundoff ub = 2−8 . Write σk2 = j pj ∥kj − µk ∥2 and σa2 = j pj (aj − µa )2 for the attention-weighted spread of keys and of a. Theorem 4.2 (Forward repair versus projection). Under these assumptions, with exact arithmetic after the cast: (a) A backward that contracts t directly,  outputs and dynamic shifting P as FA3-SBS, FP32 Psaved do, has query error Efwd = α j ξj (kj − µk ) + α j ξj µk . Over cast errors allowed by the rounding model,   X sup ∥Efwd ∥ ≥ |α| ub ∥µk ∥ |gj | − σk σa . ξ

j

Translating all keys by b leaves GQ unchanged but replaces µk by µk + b, so for any fixed GQ ̸= 0 the worst-case relative error is unbounded. P (b) GProj with r = p has query error Eproj = α j ξj (kj − µk ) and ∥Eproj ∥ ≤ |α| ub σk σa for every admissible cast; the bound does not depend on µk and is invariant under key translation. P The two bounds differ by the factor ∥µk ∥ j |gj |/(σk σa ): how far the keys sit from the origin relative to their spread (2∥µk ∥/∥k1 − k2 ∥ for two keys, at any split), which late training makes large; Figure 1 (panel 2) measures this growth on synthetic rows. This model also leaves out one effect of concentration. In the implemented kernel, the operand is formed from a BF16 saved output and reconstructed probabilities, so its row mass carries an error of order ub that does not shrink with the remaining mass, whereas the true gradient does (Section 3.3). GProj removes this row mass whatever its source. Appendix C gives the proof and the additional terms of the implemented kernel, where r is itself a BF16 operand and the center is computed from a saved output.

5

R ESULTS

5.1

ACCURACY WITH AN IDENTICAL FORWARD

GProj reduces median query/key gradient error from 219%/13% after forward repair to 0.34%/0.37% (Table 1). We evaluate eight attention inputs captured from layers 5 and 11 of the from-scratch FA3 run, and report the median relative L2 error against an FP64 reference on the same BF16 inputs; on operands from the GProj and FP32 runs, every kernel is at the BF16 floor (Table 3). FA3 is the unmodified BF16 kernel, FA3-SBS repairs its forward, GProj-Q adds only the query correction, and GProj adds both corrections. The last three produce byte-identical forward outputs, 7

Preprint

Table 1: Full-tensor relative L2 error (%) of BF16 gradients against an FP64 reference, median over eight captured attention inputs of the 450M model. FA3-SBS and the GProj rows share identical forward outputs. †,‡ See Appendix E. dQ (%)

dK (%)

dV (%)

FA3 QY-shift (source-faithful)‡ FA3-SBS GProj-Q GProj

773 1307 219 0.342 0.342

85.6 264 13.3 13.3 0.371

0.369 80.5 0.334 0.334 0.334

FP32 attention (eager)† FP32 attention (fused)†

0.384 0.385

0.371 0.371

0.165 0.165

Method

GProj

Qiu–Yao

FA3-SBS

FP32 attention

Key smoothing

(1) Training loss

2.8

0.02

2.6

0.01

2.4

0.00

2.2

(2) Pre-clipping gradient norm

loss − FP32 loss

103

FA3-SBS

GProj = FP32

20B

35B

50B

2.0

Gradient norm

3.0

Loss (nats)

FA3

1.8

102 101 100 10−1

1.6 0

10

20

30

40

50

0

Training tokens (B)

10

20

30

40

50

Training tokens (B)

Figure 3: Matched from-scratch training of a 450M model with six attention variants over 50.0B tokens; the Qiu–Yao run was stopped at 13.2B tokens (Appendix G). (1) Training loss; the inset shows the loss of FA3-SBS and GProj minus that of FP32 attention (trailing means over 2.1B tokens): GProj coincides with FP32, while FA3-SBS drifts 0.01–0.02 nats above it; (2) pre-clipping gradient norm on a log scale, with a dashed reference at 1. Curves use a light trailing average over 0.21B tokens; early loss above 3.0 is clipped, and the dotted line at 25.2B tokens marks the approximate onset of FA3’s gradient growth. so their comparison isolates the backward. FP32 attention keeps BF16 inputs and outputs but computes attention in FP32. Reference and evaluation details are in Appendix E. The query correction alone brings dQ to the level of FP32 attention, and the key pass does the same for dK. GProj does not correct dV , whose error is that of FA3-SBS. QY-shift, the released Qiu–Yao implementation (Qiu & Yao, 2026), is less accurate than FA3 on these captures (Appendix D). 5.2

T RANSLATION INVARIANCE AND ROBUSTNESS

GProj removes the leak, and each correction changes only the gradient it targets: the query correction leaves O, dK, dV and the key correction leaves O, dQ, dV byte-identical to FA3-SBS on every capture and layout tested. In a translation test, one query coordinate is zero and the matching coordinate of every key is shifted by a common value, so the exact gradient in that coordinate is zero. GProj returns essentially zero there, whereas FA3 and FA3-SBS return errors of order one, and on the four-key example of Section 3.4 GProj returns exactly zero. Beyond the captures, GProj passes all 60 cases of a synthetic suite covering common-key shifts, packed, unequal, singleton and tile-boundary layouts, and long sequences (Table 2; Appendix F). Its 8

Preprint

Table 2: Synthetic suite: relative L2 error percentages, median / maximum over the 58 test cases, excluding zero targets separately for each gradient (dQ/dK/dV : 53/51/57 nonzero targets); the two analytic calibrations are reported in Appendix F. dQ

dK

dV

0.2601 / 4794.33 0.2601 / 4794.53 0.2238 / 2.4938

0.2537 / 4.1051 0.2535 / 4.1050 0.2219 / 2.7565

0.2239 / 0.7657 0.2239 / 0.7657 0.2239 / 0.7657

Method FA3 FA3-SBS GProj

Table 3: Regime statistics at matched checkpoints, measured with exact FP32 attention on two fixed documents. Errors are relative L2 errors of dQ against an FP64 reference on the same BF16 operands, given as ratios rather than percentages (median of the two documents’ means; values above 1 mean the error exceeds the gradient). With accurate gradients (GProj, FP32) the layers stay small and FA3’s error equals cuDNN’s BF16 floor. Run (tokens)

Layers with logit scale > 103

FA3 (23.1B) FA3 (33.6B) FA3-SBS (33.6B) Key smoothing (33.6B) GProj (33.6B) FP32 attention (33.6B)

5, 11 5, 11 4, 5, 11 5, 11 none (max 26) none (max 27)

Logit scale L5 / L11

Layer 11 Layer 11 rows Layer 11 dQ error rms Q/K pmax > 0.999 FA3 / cuDNN

5,436 / 5,519 71 / 78 12,843 / 51,057 213 / 240 39,205 / 49,158 216 / 227 10,905 / 17,272 124 / 139 8/7 2/3 7/7 2/3

86% 99.4% 99.6% 74% 0% 0%

4.5 / 2.2 253 / 5.6 85 / 2.5 30 / 1.1 0.12 / 0.12 0.13 / 0.13

Table 4: Complete training-step time and peak memory; overhead is relative to the FA3 row of the same panel. Measurement protocol in Appendix E. Method

ms/update

vs. FA3

peak GiB

227.357 719.420 238.120

+0.00% +216.43% +4.73%

9.191 9.182 9.191

(2) optimized FP32 comparison FA3 231.031 FP32 attention (eager) 734.841 FP32 attention (fused) 308.276

+0.00% +218.07% +33.43%

9.191 9.182 9.191

(1) GProj comparison FA3 FP32 attention (eager) GProj

dK/dV error never rises relative to FA3-SBS; under random common-key translations its error in the zero-target coordinate is four orders of magnitude below that of FA3-SBS.

5.3

F ROM - SCRATCH TRAINING AND COST

GProj, FP32 attention and FA3-SBS finish the full 50.0B-token schedule with small gradients, while FA3 and key smoothing finish with higher loss and large gradients (Figure 3; Table 8). The released Qiu–Yao implementation stalls from the first few billion tokens, as its shift rule underflows (Appendix G). GProj ends at the same loss as FP32 attention (1.665), while FA3 ends 0.2 nats higher. FA3-SBS finishes stably but not cleanly: its backward error still drives several layers into the largelogit regime, which GProj and FP32 attention avoid (Table 3). Key smoothing delays the onset by about 3B tokens but does not prevent it, consistent with shrinking µk but not ρ. This stability comes at little cost: GProj adds 4.7% to a complete training step, against 33.4% for fused FP32 attention, and saves no extra forward state (Table 4; Appendix E). 9

Preprint

6

L IMITATIONS

Our study covers BF16 FlashAttention-3 on Hopper GPUs with head dimension 64 and a 450Mparameter model. The 4.7% overhead is measured at 4096 tokens; the GProj backward costs more than FA3’s, so the overhead grows with attention’s share of the step. GProj leaves the value gradient to the native kernel.

7

C ONCLUSION

A zero-sum conservation law ties attention’s exact gradient to key-translation invariance, and the BF16 cast breaks it. GProj restores the law after the cast with two rank-one corrections per row, bringing query/key gradient errors to the level of FP32 attention, and trains stably where FA3 degrades. Checking exact identities on the operands a kernel actually multiplies offers a practical way to audit and repair other low-precision kernels.

R EFERENCES Pierre Blanchard, Desmond J. Higham, and Nicholas J. Higham. Accurately computing the log-sumexp and softmax functions. IMA Journal of Numerical Analysis, 41(4):2311–2330, 2021. doi: 10. 1093/imanum/draa038. URL https://academic.oup.com/imajna/article/41/4/ 2311/5893596. Long Cheng, Qichen Liao, Fan Wu, Junlin Mu, Tengfei Han, Zhe Qiu, Lianqiang Li, Tianyi Liu, Fangzheng Miao, Keming Gao, Liang Wang, Zhen Zhang, and Qiande Yin. Online pseudoaverage shifting attention (PASA) for robust low-precision LLM inference: Algorithms and numerical analysis, 2025. URL https://arxiv.org/abs/2503.01873v1. Tri Dao. FlashAttention-2: Faster attention with better parallelism and work partitioning. In International Conference on Learning Representations, 2024. URL https://openreview. net/forum?id=mZn2Xyh9Ec. Tri Dao, Dan Fu, Stefano Ermon, Atri Rudra, and Christopher Ré. FlashAttention: Fast and memoryefficient exact attention with IO-awareness. In Advances in Neural Information Processing Systems, volume 35, 2022. URL https://proceedings.neurips.cc/paper/2022/ hash/67d57c32e20fd0a7a302cb81d36e40d5-Abstract-Conference.html. DeepSeek-AI. DeepSeek-V3 technical report, 2024. URL https://arxiv.org/abs/2412. 19437. Mostafa Dehghani, Josip Djolonga, Basil Mustafa, Piotr Padlewski, et al. Scaling vision transformers to 22 billion parameters. In International Conference on Machine Learning, 2023. URL https://arxiv.org/abs/2302.05442. Maxim Fishman, Brian Chmiel, Ron Banner, and Daniel Soudry. Scaling FP8 training to trilliontoken LLMs. In International Conference on Learning Representations, 2025. URL https: //arxiv.org/abs/2409.12517. FlashAttention contributors. Add the macro option for disabling FMA op in softmax calculation: Pull request 893. GitHub pull request, merged 28 March 2024, 2024. URL https://github. com/Dao-AILab/flash-attention/pull/893. Alicia Golden, Samuel Hsia, Fei Sun, Bilge Acun, Basil Hosmer, Yejin Lee, Zachary DeVito, Jeff Johnson, Gu-Yeon Wei, David Brooks, and Carole-Jean Wu. Is Flash Attention stable?, 2024. URL https://arxiv.org/abs/2405.02803v1. Aaron Grattafiori, Abhimanyu Dubey, Abhinav Jauhri, et al. The Llama 3 herd of models, 2024. URL https://arxiv.org/abs/2407.21783. Xiangming Gu, Tianyu Pang, Chao Du, Qian Liu, Fengzhuo Zhang, Cunxiao Du, Ye Wang, and Min Lin. When attention sink emerges in language models: An empirical view. In International Conference on Learning Representations, 2025. URL https://arxiv.org/abs/2410. 10781v2. 10

Preprint

Suyog Gupta, Ankur Agrawal, Kailash Gopalakrishnan, and Pritish Narayanan. Deep learning with limited numerical precision. In International Conference on Machine Learning, 2015. URL https://arxiv.org/abs/1502.02551. Alex Henry, Prudhvi Raj Dachapally, Shubham Shantaram Pawar, and Yuxuan Chen. Querykey normalization for transformers. In Findings of the Association for Computational Linguistics: EMNLP 2020, pp. 4246–4253. Association for Computational Linguistics, November 2020. doi: 10.18653/v1/2020.findings-emnlp.379. URL https://aclanthology.org/2020. findings-emnlp.379/. Byeongho Heo, Sanghyuk Chun, Seong Joon Oh, Dongyoon Han, Sangdoo Yun, Gyuwan Kim, Youngjung Uh, and Jung-Woo Ha. AdamP: Slowing down the slowdown for momentum optimizers on scale-invariant weights. In International Conference on Learning Representations, 2021. URL https://arxiv.org/abs/2006.08217v3. Robert Hu. Hardware-aware FP4 FlashAttention-4, 2026. URL https://arxiv.org/abs/ 2609.04105v1. IFM Team. Introducing K2 Horizon: Frontier performance, radically open. Blog post, Institute of Foundation Models, MBZUAI, 2026. URL https://ifm.ai/blog/k2/. Dhiraj Kalamkar, Dheevatsa Mudigere, Naveen Mellempudi, et al. A study of BFLOAT16 for deep learning training, 2019. URL https://arxiv.org/abs/1905.12322. Kimi Team. Kimi K3: Open frontier intelligence, 2026. URL https://arxiv.org/abs/ 2607.24653. Kimi Team, Yifan Bai, Yiping Bao, et al. Kimi K2: Open agentic intelligence, 2025. URL https: //arxiv.org/abs/2507.20534v1. Kimi Team, Guangyu Chen, Yu Zhang, et al. Attention residuals, 2026. URL https://arxiv. org/abs/2603.15031v1. Jingyuan Liu, Jianlin Su, Xingcheng Yao, et al. Muon is scalable for LLM training, 2025. URL https://arxiv.org/abs/2502.16982v1. Paulius Micikevicius, Sharan Narang, Jonah Alben, Gregory Diamos, Erich Elsen, David Garcia, Boris Ginsburg, Michael Houston, Oleksii Kuchaiev, Ganesh Venkatesh, and Hao Wu. Mixed precision training. In International Conference on Learning Representations, 2018. URL https://arxiv.org/abs/1710.03740. Philippe Pébay. Formulas for robust, one-pass parallel computation of covariances and arbitraryorder statistical moments. Technical Report SAND2008-6212, Sandia National Laboratories, 2008. URL https://www.osti.gov/biblio/1028931. Houwen Peng, Kan Wu, Yixuan Wei, Guoshuai Zhao, et al. FP8-LM: Training FP8 large language models, 2023. URL https://arxiv.org/abs/2310.18313. PyTorch contributors. Investigate NaNs in FlashAttention: Issue 121558. GitHub issue and technical discussion, 2024. URL https://github.com/pytorch/pytorch/issues/121558. Accessed 18 September 2026. Haiquan Qiu and Quanming Yao. Why low-precision transformer training fails: An analysis on Flash Attention. In International Conference on Learning Representations, 2026. URL https: //arxiv.org/abs/2510.04212v4. Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, and Tri Dao. FlashAttention-3: Fast and accurate attention with asynchrony and low-precision, 2024. URL https://arxiv.org/abs/2407.08608v2. Mingjie Sun, Xinlei Chen, J. Zico Kolter, and Zhuang Liu. Massive activations in large language models. In First Conference on Language Modeling, 2024. URL https://arxiv.org/ abs/2402.17762v2. 11

Preprint

Mitchell Wortsman, Peter J. Liu, Lechao Xiao, Katie Everett, et al. Small-scale proxies for largescale transformer training instabilities. In International Conference on Learning Representations, 2024. URL https://arxiv.org/abs/2309.14322. Guangxuan Xiao, Yuandong Tian, Beidi Chen, Song Han, and Mike Lewis. Efficient streaming language models with attention sinks. In International Conference on Learning Representations, 2024. URL https://arxiv.org/abs/2309.17453v4. Jianlin Yu, Jing Lin, Linghui Kong, Aiyue Chen, Weiyi Sun, Chenyu Zeng, Wangli Lan, Jinxi Li, Zhuo Zheng, Ziyang Yue, Danning Ke, Fei Yi, Tianchi Hu, Yuan Ding, Yiwu Yao, and Junsong Wang. MXAttention: Data-free optimal scaling and pre-normalization quantization for MXFP4 attention, 2026. URL https://arxiv.org/abs/2607.24377. Pedram Zamirai, Jian Zhang, Christopher R. Aberger, and Christopher De Sa. Revisiting BFloat16 training, 2020. URL https://arxiv.org/abs/2010.06192. Shuangfei Zhai, Tatiana Likhomanenko, Etai Littwin, Dan Busbridge, Jason Ramapuram, Yizhe Zhang, Jiatao Gu, and Joshua M. Susskind. Stabilizing transformer training by preventing attention entropy collapse. In Proceedings of the 40th International Conference on Machine Learning, volume 202 of Proceedings of Machine Learning Research, pp. 40770–40803. PMLR, 2023. URL https://proceedings.mlr.press/v202/zhai23a.html. Jintao Zhang, Haofeng Huang, Pengle Zhang, Jia Wei, Jun Zhu, and Jianfei Chen. SageAttention2: Efficient attention with thorough outlier smoothing and per-thread INT4 quantization, 2025a. URL https://arxiv.org/abs/2411.10958v7. Jintao Zhang, Jia Wei, Haofeng Huang, Pengle Zhang, Jun Zhu, and Jianfei Chen. SageAttention: Accurate 8-bit attention for plug-and-play inference acceleration, 2025b. URL https: //arxiv.org/abs/2410.02367v9. Jintao Zhang, Marco Chen, Haoxu Wang, Kai Jiang, Ion Stoica, Joseph E. Gonzalez, Jianfei Chen, and Jun Zhu. SageBwd: A trainable low-bit attention, 2026. URL https://arxiv.org/ abs/2603.02170v1.

12

Preprint

Appendix roadmap. Appendices A–C give the cast witness, implemented arithmetic and proofs; Appendix D defines the comparison methods; Appendix E gives the accuracy and cost protocols; Appendix F the synthetic suite; Appendix G the training runs and the regime audit. The remaining appendices give the model configuration, the full localization of the FA3 failure (saved-output channel, fixed-forward interventions, source intervention, saved-state crosses, backend comparison) and a plain-Transformer control.

A

E XACT T RANSLATION S TRUCTURE AND THE C AST W ITNESS

We use gauge for a common translation of keys or values that leaves the exact local VJP unchanged at fixed incoming u, and the row notation of Section 3.2 throughout. A.1

E XACT TRANSLATIONS AND NUMERICAL ROW- MASS LEAKAGE

Proof of Proposition 3.1. For a common key translation kj 7→ kj +b, all logits in one row change by the same αq ⊤ b and the exact softmax probabilities are unchanged. For a common value translation P vjP 7→ vj +c, every aj changes by u⊤ c, which cancels in a − p a j ℓ ℓ . Thus g, and the contributions ℓ P α j gj kj , αgj q and pj u, are unchanged because j gj = 0. Expanding the centered covariance P gives GQ = α j pj (kj − µk )(aj − µa ). These identities also underlie key smoothing in SageBwd (Zhang et al., 2026, Section 6); the value translation changes the forward output by c, so u must be held fixed. They single out the direction, the common key offset, to which a numerical backward can respond spuriously. Key smoothing removes a common mean from the keys before quantization, which changes the key operand but not the row sum of the score gradient formed afterwards. GProj instead corrects the cast score operand itself. Row-mass leakage. Proposition A.1 (Exact leakage P decomposition). Let t be any numerical score-gradient operand, let e = t − g, and let ρ = j tj . Before any additional contraction rounding, Equation (4) gives the decomposition. For fixed t, a common-key translation P b changes its contraction by exactly αρb. More generally, subtracting λr from t, where m = j rj ̸= 0, leaves row mass ρ − λm. Thus λ = ρ/m eliminates this mass in real arithmetic, whereas λ = ρ leaves ρ(1 − m). P P P P Proof. Since j gj = 0, ρ = j tj = j ej . Adding and subtracting µk inside α j ej kj gives P P α j ej kj = α j ej (kj − µk ) + αρµk , which is Equation (4). For fixed t, translating every key by b adds αρb to the contraction, and subtracting λr from t leaves row sum ρ − λm. The last statement is why GProj divides by the actual mass m of the BF16 probabilities it multiplies: the mass of any other probability representation does not cancel ρ. An exact −24 witness. A four-key row with an exact forward exhibits the leak. Take four equally likely keys, α = 1/8, q = 0, u = (1, 1), η = 3/1024, and kj = (65536, j − 1),

(v1 , v2 , v3 , v4 ) = ((4, 4η), (−4, 0), (−4η, 0), (0, 0)).

(8)

Every input is exactly BF16-representable. Here pj = 1/4 and a = (4 + 4η, −4, −4η, 0) has mean zero, so the output is o = (−η, η), u⊤ o = 0 exactly, and g = (1 + η, −1, −η, 0). BF16 rounds 1 + 3/1024 to one and represents the other three entries exactly, so t = (1, −1, −η, 0) with ρ = −η. The true first coordinate of GQ is zero; contracting the cast operand gives −(1/8) 65536 (3/1024) = −24. An exact saved output does not repair this cast channel. Zero padding embeds the construction in d = dv = 64. 13

Preprint

B

I MPLEMENTED A RITHMETIC AND ROUNDING

The GProj kernel targets Hopper GPUs with BF16, head dimension 64 and GQA, in dense and packed causal layouts. Matrix operands, outputs and returned gradients are BF16; accumulators, row statistics and scalar corrections are FP32. GProj consists of the native backward with a query tile of 96, followed by a separate dK correction pass. The projection acts on the BF16 operands after they are formed. Let pb denote the probabilities reconstructed by the native backward and let b aj and δb denote its computed value product and savedoutput reduction. The operands consumed by its BF16 matrix multiplications are X X  b , rj = castBF16 (b pj ), tj = castBF16 fl32 [b pj (b aj − δ)] m= rj , ρ = tj . (9) j

j

Here fl32 evaluates the bracketed operations in FP32, and m, ρ are exact sums of the formed operands, which the kernel approximates in FP32. The attention scale α is applied later, and in general m ̸= 1 and r equals neither p nor pb. Two passes. The first pass forms t and r, measures their row masses, accumulates tK and rK, and writes the corrected dQ while keeping the FP32 key accumulator. The second pass revisits the query/key tiles, reconstructs the probabilities and corrects that accumulator before dK is cast to BF16. Algorithm 1 gives the tiled pseudocode; the equations below fix the rounding and the placement of the scale. Forward and saved state. For computed unscaled logits xij and their row maximum xi,max , the FA3-SBS forward explicitly rounds the subtraction: zij = exp232 (fl32 [sub32 (xij , xi,max )(α log2 e)]) .

(10)

Everything else in the forward is native, and it saves BF16 O and the FP32 log-sum-exp (LSE) as usual. Query correction. After the native preprocessing computes δbi from the saved Oi and Ui , the backward forms t and r as in Equation (9), measures their masses and accumulates two contractions: bi = acc32 (ti K), A

bi = acc32 (ri K), B

bQ = divQ (b λ ρi , m b i ), i

(11)

where acc32 is a BF16 matrix product accumulated in FP32, hats denote computed quantities and divQ is the compiled FP32 division. After the last key tile,   c = castBF16 fl32 [α(A bQ B bi )] . bi − λ dq (12) i i The scale is applied once, after the correction, and a row with zero computed mass receives no correction. Key correction. The same projected operand gives the ideal key gradient Gproj = α[t⊤ Q − K ⊤ r (λ ⊙ Q)], where λ ⊙ Q scales each query row. Since λ is known only after a row has seen every key, the key correction needs its own pass over the tiles. It reuses the stored row statistics but divides bK = divK (b again, in Triton FP32 arithmetic with the same zero-mass rule, so λ ρi , m b i ) need not i Q b bitwise. It computes match λ i   ⊤ bK qi ), dK c = castBF16 fl32 [A b reij = castBF16 (e pij ), zi = castBF16 (λ − α acc (e r z)] , K 32 i (13) bK is the native key accumulator, already scaled and kept in FP32, so the correction happens where A bK qi adds before the final BF16 cast. Because the reconstructed re can differ from r and casting λ i error, dQ and dK approximate the same ideal projection without sharing a bitwise identical score operand. The key pass leaves O, dQ and dV untouched. Cost. Relative to the native backward, GProj adds one rK contraction, a query-shaped FP32 workspace, two FP32 row-statistic arrays and the second pass. It saves no extra forward state and stores no quadratic matrix; dV is not projected. 14

Preprint

C

G AUGE P ROJECTION : P ROOFS

All statements concern exact attention at fixed represented operands and a fixed incoming derivative. C.1

P ROJECTION AND THE ERROR THAT REMAINS

With m > 0, Πr = I − r1⊤ /m obeys Π2r = I − 2r1⊤ /m + r(1⊤ r)1⊤ /m2 = Πr ,

1⊤ Πr = 0.

It acts as the identity on vectors of zero row sum. Since g is such a vector, Πr t − g = Πr (t − g), and contracting the latter gives P X X X X j ej α (Πr e)j kj = α ej kj − α rj kj = α ej (kj − µr ). m j j j j The same algebra gives translation invariance at fixed t, r and completes the proof of Proposition 4.1. The projection removes only the row mass of the error; an error that already sums to zero is left intact. Centering before the cast. The remaining error depends on how the operand is centered before it is cast. P Lemma C.1 (Conditional centered-cast bound). Suppose pj > 0, j pj = 1, and the exact pre-cast operand is xj = pj (aj − β). Suppose its cast tj = xj + ξj satisfies |ξj | ≤ ub |xj |. If all subsequent projection and contraction arithmetic is exact and the correction uses r = p, then α

X

(Πp t)j kj − GQ

≤ |α|ub σk

j

where σk2 =

2 2 j pj ∥kj − µk ∥2 and σa =

P

p σa2 + (µa − β)2 ,

(14)

2 2 j pj (aj − µa ) .

P

Proof. P The exact projection gives Πp x = p ⊙ (a − µa ) = g, irrespective of β. Thus the query error is α j ξj (kj − µk ). Weighted Cauchy–Schwarz gives X j

1/2  1/2  X X ξj2  σk ≤ ub  pj (aj − β)2  σk . ξj (kj − µk ) ≤  p 2 j j j

The final sum is σa2 + (µa − β)2 . The error grows with |µa − β|, so an accurate center before the cast still matters: projection cannot restore value differences lost in the cast. C.2

P ROOF OF T HEOREM 4.2

P With a perfect forward, x = g and e = t − g = ξ. Part (a): contracting t gives α j tj kj − P GQ = α j ξj kj , and adding and subtracting µk yields the stated decomposition, exactly as in Equation (4).PFor the lower bound, choose the admissible errors ξj = ub |gj |, all of one sign, so that P j ξj = ub j |gj |. By the triangle inequality, X X ∥Efwd ∥ ≥ |α| ub ∥µk ∥ |gj | − |α| ξj (kj − µk ) , j

j

P P and weighted Cauchy–Schwarz bounds the last norm by ( j ξj2 /pj )1/2 σk ≤ ub ( j pj (aj − µa )2 )1/2 σk = ub σa σk , using |gj | = pj |aj − µa |. A common key translation by b leaves p, a, g and GQ unchanged and replaces µk with µk + b while kj − µk is unchanged; taking ∥b∥ → ∞ along any direction makes the bound, and hence the worst-case relative error, unbounded. Part (b) 15

Preprint

is Lemma C.1 with β = µa : the projection with r = p maps x to g, the error is α and the same Cauchy–Schwarz step gives |α|ub σk σa . Neither µk nor b enters.

P

j ξj (kj − µk ),

Dynamic shifting (Qiu & Yao, 2026) and FP32 saved outputs change only how p and µa are formed, so with a perfect forward they fall under part (a). In the implemented kernel, r is a separate BF16 cast with m ̸= 1, the center comes from a saved output and the contractions round; the analysis below accounts for these terms. Synthetic offset sweep. Panel 2 of Figure 1 checks the theorem numerically. Each of 256 rows has L = 1024 keys kj = zj + b in d = 64 dimensions, with zj ∼ N (0, I) and a common translation b along a random direction, ∥b∥ ∈ [1, 104 ]. Queries give logits with standard deviation 3, and values and the upstream derivative are standard normal. The forward and the score gradient g are exact in FP64; the only rounding is t = BF16(g). The forward-repair curve contracts t directly, GProj subtracts (ρ/m) r with r = BF16(p), and both contractions are evaluated in FP64. The translation leaves p, g and GQ unchanged; the horizontal axis is the median of ∥µk ∥/σk at each offset. Logit standard deviations of 1 and 6 change the forward-repair curve by less than 10% and move the GProj floor between 0.12% and 0.17%. C.3

ROUNDING IN THE IMPLEMENTED QUERY CORRECTION

The kernel adds accumulation and division P errors to the ideal To separate them, P projection. P P define exact sums of the formed operands by A = j tj kj , B = j rj kj , ρ = j tj and m = j rj > 0. Write the actual FP32 accumulator/statistic errors as bQ = λ + eQ , λ = ρ/m. b = A + eA , B b = B + eB , ρb = ρ + eρ , m A b = m + em , λ λ

Let eF include the rounding of the final multiply/subtract, scale and BF16 conversion. When all quantities in the decomposition are finite, the implemented query error satisfies the exact decomposition X bQ eB − αeQ B + eF . b Q − GQ = α G (tj − gj )(kj − µr ) + αeA − αλ (15) λ j

The first term is the operand error that survives the projection; the others are the accumulation, ratio and final rounding errors, and the norm of the total is at most the sum of their norms. bQ = ρb/m If |em | < m and λ b + eQ div , then |eQ λ| ≤

|eρ | + |λ| |em | + |eQ div |. m − |em |

(16)

bQ r This follows by subtracting ρ/m from (ρ + eρ )/(m + em ). The hypothetical coefficients t − λ Q have row mass −meλ even before the contraction is rounded. A row with zero computed mass receives no correction. C.4

ROUNDING IN THE IMPLEMENTED KEY CORRECTION

The key pass adds two further sources of error: it reconstructs the probabilities and casts the scaled queries. Index every query/head row by i, and sum only over rows legally attending to a key in its KV head. The ideal projected key gradient is X Gproj = α (tij − λi rij )qi . K,j i

It contracts the same projected operand as Equation (6); GQA only enlarges the set of contributing query heads. The implemented second pass instead has a reconstructed BF16 probability reij and bK qi ). Put ez,i = zi − λ bK qi and eK = λ bK − λi . The key scaled-query operand zi = castBF16 (λ i i i λ,i bK = ρbi /m ratio is recomputed from the shared ρbi , m b i ; write λ b i + eK . This division error need not i

div,i

equal the native query division error. For |em,i | < mi , |eK λ,i | ≤

|eρ,i | + |λi | |em,i | + |eK div,i |. mi − |em,i | 16

Preprint

Method FA3 FA3-SBS

Forward

BF16 FA3 Subtract-before-scale FA3 Key smoothing FA3 on mean-subtracted keys QY-shift Author-source dynamic shift GProj Subtract-before-scale FA3 FP32 attention FP32 attention

Backward

Saved output

Native mixed precision Native mixed precision

BF16 BF16

Native mixed precision

BF16

Author-source BF16 arithmetic

BF16

Dense query/key projection

BF16

FP32 attention

See below

Table 5: Primary method definitions. All methods expose BF16 output and gradient tensors to the model. GProj retains BF16 matrix products with FP32 accumulation and corrections. Eager FP32 recomputes its output; fused FP32 saves an internal FP32 output. Keeping every operand error, the exact product mismatch is reij zi − rij λi qi = (e rij − rij )λi qi + reij eK eij ez,i . λ,i qi + r

(17) P

If eK,j denotes the error of the already scaled native FP32P accumulator relative to α i tij qi , eC,j denotes the FP32 correction-contraction error relative to i reij zi , and eF,K,j denotes final arithmetic and output-cast error, then X b K,j − Gproj = eK,j − αeC,j − α G (e rij zi − rij λi qi ) + eF,K,j . K,j i

Recomputation and the extra BF16 cast are why dQ and dK need not share an identical effective score operand. Correcting the FP32 accumulator, rather than the rounded BF16 output, avoids one more rounding step. The value gradient is not projected.

D

C OMPARISON M ETHODS AND A RITHMETIC C ONTRACTS

Methods. Table 5 summarizes the methods, which differ only in the attention computation. FA3 is unmodified FlashAttention-3. FA3-SBS changes only its forward softmax, subtracting the row maximum before scaling (Appendix I.3). GProj adds the query and key corrections to the FA3-SBS forward and, like FA3, saves BF16 O and the FP32 LSE. FP32 attention computes attention in FP32 behind the model’s BF16 interface; we use an eager implementation, which recomputes from Q, K, V , and a fused one, which saves an FP32 output. QY-shift is the dynamic-shift method of Qiu and Yao (Qiu & Yao, 2026, Section 4), with the arithmetic of their released BF16 source (STABLE=1). For a current tile of scaled, masked logits, let r P be its row maximum and c = j 1{sj ≥ flBF16 (r − ϵ)}. The source computes  2r, r > 0 and c > 1, b = 0, r < 0 and c > 1, mnew = max(mold , b), (18)  r, otherwise, with ϵ = 10−3 . Preserving the rounded subtraction in the count predicate matters in BF16. The previous numerator and denominator are rescaled by exp(mold − mnew ) before the online update. The implementation retains 512 × 512 tiles, the authors’ 10−10 tile-sum clamp, BF16 intermediate state and gradient accumulation, and BF16 saved O and LSE. The source keeps BF16 precision between operations as well, not only at its interface. Writing dP = U V ⊤ and D for the computed row reduction of U ⊙ O, the source rounds that elementwise product before its row reduction, forms the scaled score operand as (P α) ⊙ (dP − D) in BF16, and updates BF16 gradient buffers after each tile. GProj instead retains FP32 gradient accumulators and applies α after the query correction. Our C++/CUDA port matches the author code byte for byte (Appendix E). 17

Preprint

Table 6: Per-capture relative L2 errors (%) of GProj and fused FP32 attention against FP64 references; Table 7 identifies the captures. Capture

GProj dQ

dK

dV

FP32 dQ

dK

dV

0 1 2 3 4 5 6 7

0.3223 0.2504 0.2669 0.2668 1.2604 0.4158 0.9490 0.3612

0.2942 0.2397 0.3183 0.2855 1.8577 0.5350 1.3361 0.4240

0.1699 0.1924 0.1734 0.1966 1.0888 0.4715 1.1936 0.5267

0.4203 0.2245 0.2799 0.4146 1.0951 0.3552 0.9065 0.3187

0.3290 0.2216 0.3411 0.3760 2.7637 0.3657 2.5611 0.3764

0.1592 0.1648 0.1603 0.1639 0.1771 0.1648 0.1745 0.1673

Table 7: Identity key for Table 6. Layers are zero-based; P is a PG19 book and D a ProofPile document. Capture

Checkpoint

Layer

Document

0 1 2 3 4 5 6 7

5500 5500 5500 5500 8000 8000 8000 8000

11 5 11 5 11 5 11 5

P P D D P P D D

Ablations. The ablations build GProj up from FA3-SBS, which shares its forward and keeps the native backward. GProj-Q adds the query correction, together with the new backward tiling, and GProj adds the key pass; GProj always denotes the complete method. Because all three share one forward, their differences are differences of the backward, and the key pass is measured before and after correction on the same execution.

E

ACCURACY AND C OST: S ETUP AND A DDITIONAL R ESULTS

Captured inputs. We evaluate all kernels on eight attention inputs captured from the from-scratch FA3 run: layers 5 and 11 at 23.1B and 33.6B tokens (checkpoints 5500 and 8000), on one PG19 book and one ProofPile document (Table 7). Each capture holds BF16 Q, K, V and the incoming derivative U for one causal sequence of 4096 tokens with 16 query heads, four KV heads, head dimension 64 and scale 1/8. Accuracy can therefore be reproduced from the captures alone, without checkpoints or the trainer. Reference and metric. The reference is the exact attention VJP at the represented BF16 inputs, computed in FP64 with centered scores and values and with all GQA contributions summed. We 64 b report 100∥G−G ∥2 /∥G64 ∥2 for each full gradient tensor, and medians and maxima over the eight captures. The FP32 attention study evaluates the same target with an independent FP64 implementation (explicit softmax Jacobian, 128-query chunks), and QY-shift is scored against a separately implemented centered FP64 reference. Table 6 lists the errors of every capture. Backward-only comparison. GProj uses the FA3-SBS forward, and on every capture its output is byte-identical to that of FA3-SBS, so comparing FA3-SBS with GProj compares backward passes only. The GProj backward targets Hopper’s SM90a instructions with head dimension 64 and a query tile of 96, against 128 in FA3. The query correction alone (GProj-Q) leaves O, dK and dV unchanged in 96 checks over the captures, 20 additional layouts and lengths, and four keytranslation cases, so its median dK error stays at 13.2651%. The key pass then corrects the FP32 key accumulator before the final BF16 cast, using BF16 probabilities and BF16 scaled queries with FP32 accumulation. On the same execution it lowers the median dK error from 13.2651% to 0.3712% while leaving O, dQ and dV unchanged. 18

Preprint

QY-shift. QY-shift (Appendix D) ports the released Qiu–Yao BF16 code without guards or fallbacks, moving the tile loops to C++ with CUDA pointwise kernels and cuBLAS matrix products. On nine synthetic cases (tile boundaries, repeated and large logits, noncontiguous tensors, packed segments, GQA) and on the eight captures, its output, LSE and gradients match the author code byte for byte. Autocast and compilation are disabled. Fused FP32 baseline. We built a fused FP32 kernel. It decodes the BF16 inputs, computes in FP32 SIMT/FMA arithmetic without TF32, and returns BF16 outputs and gradients. It saves FP32 O and three FP32 row statistics: the LSE, the row maximum and the inverse normalizer, the last two keeping normalization accurate when large logits leave the LSE short of resolution. At 64 sequences per GPU the saved state is 1.046875 GiB, and no N 2 matrix is stored. Timing. Table 4 measures complete training steps of the 450M model on one H200 with batch one and 4096 tokens, under the training configuration: FP32 master weights, BF16 compute, the Muon/AdamW optimizer, clipping at one and activation checkpointing. Each step runs 48 attention forward calls, including recomputation, and 24 backward calls; data loading and communication are excluded. Each method runs on three workers, each timing 15 updates after warmup, and we report the median over workers of the mean step time. The two panels are separate jobs on different nodes. Peak allocated memory is 9.191 GiB for both FA3 and GProj and 9.182 GiB for eager FP32.

F

S YNTHETIC N UMERICAL S UITE

The synthetic suite contains 58 cases that stress the layouts and numerical regimes an attention kernel must handle, together with two analytic four-key calibrations. The cases cover packed segments, unequal lengths, empty causal rows, singleton and tile-boundary rows, batch size two, non-causal attention, common key and value shifts, constant values, sharp attention, zero and tiny upstream gradients, and four 4096-token sequences, all in BF16 with head dimension 64 and GQA. The suite takes 65 seconds on one H200 and evaluates O, dQ, dK, dV for FA3, FA3-SBS and GProj. Table 2 summarizes the errors. References. Each case is scored against two independent FP64 references, a centered VJP and an explicit softmax-Jacobian VJP. They agree to within 10−10 + 10−9 times the target norm on all 60 cases, the worst case using 0.5154% of that tolerance (6.30643 × 10−10 ), and eleven small cases further agree with autograd to 9.99 × 10−16 . Outcome. A case fails if its relative error exceeds 5%, or if its absolute error exceeds 10−6 where the target tensor is zero. GProj passes every case and never increases the dK or dV error; its query error exceeds that of FA3-SBS in the three cases below. Its largest relative dQ and dK errors in Table 2 both occur with tiny upstream gradients. Case

FA3-SBS (%)

GProj (%)

0.227716 2.448365 0.194175

0.227779 2.493797 0.582524

015: unequal lengths 048: tiny upstream 058: analytic calibration

Exact zeros. Where the exact answer is zero, GProj stays close to it. The singleton case 000 has exactly zero dQ and a dK residual of 5.94918 × 10−10 , and the constant-value cases 028 and 041 have dQ/dK residuals of 2.02560 × 10−9 /2.72331 × 10−9 and 2.43146 × 10−9 /3.05945 × 10−9 . Under random common-key shifts the zero-target coordinate reaches at most 0.001953125 (case 053) and 0.01171875 (case 057), against 432 for FA3-SBS in the latter. Operand masses and structure. Over 295,008 query-head rows the BF16 probability mass m is always positive and finite, with minimum 0.9967041015625 and |m−1| at most 0.00341796875. All accumulator mappings and second-pass checks pass, GProj’s O and dV are byte-identical to those of FA3-SBS, and repeated runs of one packed and one 4096-token case are bitwise reproducible. 19

Preprint

Tokens (B)

FA3

FA3-SBS

0.8 4.624 / 0.827 12.6 1.870 / 0.406 23.1 1.783 / 0.358 29.4 1.909 / 3.871 33.6 2.122 / 375.508 37.7 2.107 / 446.931 41.9 1.949 / 70.725 46.1 1.882 / 110.166 50.0 1.871 / 39.015 Final 0.84B Norm >10 at

GProj

4.635 / 0.790 1.871 / 0.360 1.779 / 0.375 1.783 / 0.325 1.780 / 0.364 1.753 / 0.279 1.749 / 0.457 1.679 / 0.210 1.684 / 0.125

1.866 / 30.705 29.9B

Key smoothing FP32 attention

4.623 / 0.918 4.618 / 0.804 1.869 / 0.312 1.870 / 0.373 1.773 / 0.403 1.776 / 0.329 1.773 / 0.385 1.790 / 0.355 1.769 / 0.396 1.937 / 7.381 1.740 / 0.324 1.932 / 4.789 1.728 / 0.335 2.026 / 138.966 1.664 / 0.142 1.912 / 138.389 1.670 / 0.076 1.907 / 287.047

4.625 / 0.852 1.869 / 0.372 1.772 / 0.359 1.773 / 0.288 1.769 / 0.354 1.740 / 0.375 1.728 / 0.270 1.664 / 0.168 1.670 / 0.095

1.679 / 0.112 1.665 / 0.087 1.903 / 233.384 never never 33.2B

1.665 / 0.097 never

Table 8: Matched from-scratch training (seed 42, same data order): all-rank mean training loss and mean pre-clipping gradient norm per 10-step logging interval. Cells show loss / gnorm. All five arms complete 50.0B tokens. The Qiu–Yao arm, stopped at 13.2B tokens, is reported in Appendix G (Figure 4). “Final 0.84B” gives mean loss / median gnorm over the 20 records in the final 0.84B tokens; “Norm >10 at” gives the tokens at which the gradient norm first exceeds 10. (1) Training loss

(2) Gradient norm

FA3 GProj Qiu–Yao

Pre-clipping norm

Loss (nats)

3.0 2.8 2.6 2.4 2.2 2.0

6 × 10

(3) Shift rule on checkpoints

0

Qiu–Yao: rows fully clamped

60

Qiu–Yao: 2r-shift tiles

0

4 × 10 3 × 100

% (worst layer)

3.2

2 × 100 100

6 × 10−1 −1

4 × 10 3 × 10−1

FA3 weights, same rule

50 40 30 20 10

1.8

0 0.0

2.5

5.0

7.5

10.0

12.5

Training tokens (B)

0.0

2.5

5.0

7.5

10.0

Training tokens (B)

12.5

0.0

2.5

5.0

7.5

10.0

12.5

Training tokens (B)

Figure 4: The Qiu–Yao baseline arm. (1) All-rank mean training loss and (2) pre-clipping gradient norm of the Qiu–Yao, FA3 and GProj arms over the Qiu–Yao arm’s 13.2B tokens, all-rank means without smoothing; the dotted line marks the spike at 6.3B tokens (step 1500). (3) Diagnosis on the Qiu–Yao checkpoints (eight fixed documents, every sixteenth query row): the worst-layer percentage of query rows whose every 512-key tile sum falls below the source’s 10−10 clamp, and of legal row– tile pairs taking the 2r shift branch; the FA3 arm’s checkpoints under the same rule are shown for comparison. The text of this appendix gives the mechanism.

G

T RAINING S TUDY

Setup. All matched runs share one recipe: 16 H200 GPUs, seed 42, the same data order, and 11,930 updates of 4,194,304 tokens (50.0B tokens). They differ only in the attention computation. Key smoothing subtracts the per-document, per-KV-head mean key before calling unmodified FA3 (FP32 accumulation, BF16 result); softmax is invariant to this shift in exact arithmetic, and autograd applies the matching correction to dK. Key smoothing keeps the FA3 forward; FA3-SBS, which repairs only the forward, still enters the same large-logit regime (Table 3). The Qiu–Yao run was stopped at 13.2B tokens, for reasons given below. Table 8 and Figure 4 report the all-rank mean loss and mean pre-clipping gradient norm of each ten-update logging interval, without smoothing; when a run resumed from a checkpoint, the resumed records replace the superseded ones. Large-logit regime. Each run’s checkpoint at 33.6B tokens (and FA3’s at 23.1B) is evaluated with exact FP32 attention on two fixed documents. Table 3 reports the regime statistics and the error that unmodified FA3 and cuDNN would make on those operands. Qiu–Yao run. The Qiu–Yao run falls behind FA3 from about 2.5B tokens, spikes at 6.2–6.8B tokens (gradient norm 6.4) and then stalls near a loss of 2.28–2.30 while FA3 keeps improving, 20

Preprint

trailing it by 0.31 nats at 8.4–8.9B tokens (Figure 4). The cause is the dynamic shift rule, not the attention backward. When a tile’s row maximum r > 0 has a near-tie (c > 1), the rule shifts by 2r, so every weight is at most exp(s − 2r) ≤ exp(−r). Once logits are large, near-ties are common, because the BF16 spacing at |s| ≥ 128 is 1, and exp(−r) leaves the FP32 normal range for r > 87. The tile sum then falls below the source’s 10−10 clamp, and the row’s output is dominated by rounding. Replaying the author arithmetic on the saved checkpoints (eight documents, every sixteenth query row) shows this failure growing with training. At step 1000 (4.2B tokens) the 2r branch is taken on 12–29% of row–tile pairs in layers 4–23, 5–25% of query rows have every tile clamped, and the attention output differs from FP32 attention on the same operands by a median relative error of 0.5– 0.9. At step 1500 (6.3B) the branch fraction reaches 18–42%, up to 36% of rows are fully clamped and row maxima reach r ≈ 5,300; by step 3000 (12.6B) about half of the rows in layers 15–23 are fully clamped and the output is uncorrelated with FP32 attention (median relative error ≈ 1.0). The corrupted output in turn feeds logit growth: the run’s 99th-percentile row maxima are 5–100× those of FA3 at the same steps, where the same rule would fully clamp at most 11% of rows. The weights also adapt to the clamped outputs: substituting exact FP32 attention does not improve the fixed-document NLL from step 1000 to 2000 (3.32 to 3.30, and 3.31 at step 3000), whereas the FA3 arm improves from 3.16 to 2.92 over the same steps; at step 1500 the substitute is even worse than the run’s own arithmetic (3.61 versus 3.51).

H

M ODEL AND T RAINING C ONFIGURATION

The 450M model in the main text is a Full-AttnRes transformer (Table 9), using learned softmax aggregation over earlier sublayer outputs (Kimi Team et al., 2026). Its shared-read configuration uses one residual-depth read jointly for token-attention queries, keys and values. All training runs, captured operands and frozen-state diagnostics use this configuration, and every run is trained from scratch. Setting

Configuration

Parameters Transformer Attention Vocabulary Sequence length Attention backend Normalization controls Parameter precision Matrix optimizer Other parameters Schedule

449,465,344; referred to as 450M in the main text 24 layers, hidden width 1024 16 query heads, four key/value heads, head dimension 64 64,256; untied input and output embeddings 4096 tokens Native BF16 FlashAttention-3 Query–key normalization and softcapping disabled BF16 compute parameters, FP32 master parameters Muon, learning rate 0.01, momentum 0.99 AdamW, learning rate 0.001, β = (0.9, 0.95) Warmup 300 updates, constant, cosine decay over the final 9,544 updates to 1% Zero weight decay; no dropout Global norm threshold 1 Two nodes, 16 ranks, 64 sequences per rank; global batch 1024 sequences Enabled layerwise 11,930 optimizer updates (50.0B tokens), seed 42, deterministic

Regularization Gradient clipping Execution Activation recomputation Training length

Table 9: Configuration of the 450M model and of every training run. The comparison arms change only the attention computation. The plain-transformer control has its own configuration (Appendix K).

I

L OCALIZING THE N UMERICAL FAILURE

All analyses use checkpoints of the from-scratch FA3 run. 21

Preprint

After saved O Ckpt. Layer dQ rel. err. Resid. frac. Err. cos. Saved-O red. rel. err.

cos.

5500 5500 8000 8000

0.419 0.366 0.686 0.156

5 11 5 11

6.517 6.120 17.889 411.104

0.0055 0.0060 0.0067 0.0197

1.0000 1.0000 1.0000 0.9999

2.98× 2.45× 15.73× 72.02×

2.165 2.496 1.055 5.988

Table 10: Confirmation on 127 documents per checkpoint/layer, with 128 strided query rows and all 4096 keys. Entries are per-document medians; relative errors are ratios, not percentages. Resid. frac. and Err. cos. compare Equation (19) with observed native error. Saved-O red. is the median paired reduction in absolute dQ error after replacing only saved O by an FP64-recomputed output rounded to BF16; the last two columns give the remaining relative dQ error and the cosine between the computed and exact dQ. I.1

T HE SAVED - OUTPUT CHANNEL PREDICTS THE LOCAL ERROR P FA3’s backward forms the reduction δi = j pij u⊤ i vj from the saved output, a known sensitivity of low-precision attention (Qiu & Yao, 2026). With probabilities and u⊤ i vj exact, an error in this reduction produces the query-gradient error X Eδ,i = −α(δbi − δi⋆ ) p⋆ij kj , (19) j

where δ ⋆ and p⋆ are FP64 reference values and δbi = dot32 (ui , obi ) is computed from the native saved output. We test this prediction on native BF16 Q, K, V, O and incoming derivatives U captured at layers 5 and 11, after checking that rerunning the native forward reproduces the saved output c − dq ⋆ through their cosine and the residual bitwise, and compare it with the observed error Ei = dq i i fraction ∥E − Eδ ∥2 /∥E∥2 . The analysis uses 128 strided query rows against all 4096 keys, on an exploratory set of eight documents and a confirmation set of 127 documents at checkpoints 5500 and 8000 (23.1B and 33.6B tokens); the kernel comparisons elsewhere use all 4096 rows. One document shows the effect. For a PG19 book at checkpoint 8000, layer 11, the native ∥dQ∥2 is 4.30 × 10−3 against a reference norm of 4.27 × 10−5 , a relative error of 101, whereas the same VJP computed in FP32 has relative error 0.015. The prediction has cosine 0.9999 with the native error and leaves a residual fraction of 0.020. The forward output, by contrast, has relative error only 0.00193. Table 10 shows that the same holds across 127 documents. At checkpoint 8000, layer 11, the median residual fraction is 0.0197, and replacing only the saved O reduces the absolute dQ error by a median factor of 72.02; the earlier checkpoint shows the same channel with smaller amplification. I.2

F IXED - FORWARD INTERVENTIONS REACH THE FULL MODEL

We rerun the full backward on eight documents at checkpoints 6500 and 8000 (27.3B and 33.6B tokens) and change only the attention backward of layers 5 and 11: either the native backward receives a correct saved output (recomputed in FP64 and rounded to BF16), or the exact attention VJP is computed in FP32. All 24 attention outputs and the loss stay bitwise identical, but a change at layer 11 propagates to the derivative entering layer 5, so these are whole-model effects. Both repairs remove almost all of the excess gradient (Table 11). At checkpoint 8000, correcting the saved output lowers the median norm from 5353.856 to 19.746, and the FP32 VJP gives 18.185. At checkpoint 6500, before the training gradient norm first exceeds 10, the native norm is already 36.653, and the same repairs bring it to about 2.3–2.5. Repairing only the first query row, by contrast, changes almost nothing. That row has a single legal key, so p00 = 1, o0 = v0 and dq0 = 0 exactly, and native captures violate these identities (Section I.3); yet correcting its saved output or supplying its exact VJP barely changes the model gradient. A violated exact identity can reveal a kernel defect without locating the rows that dominate training. 22

Preprint

Backward intervention at layers 5 and 11 Native Replace saved O, first row only Exact first-row VJP only Replace saved O, all rows FP32 VJP, all rows

6500

8000

36.653 36.402 36.402 2.450 2.258

5353.856 5353.857 5353.857 19.746 18.185

Table 11: Median full-model parameter-gradient norm over eight documents. All 24 attention outputs and the scalar loss are bitwise identical across modes at each document/checkpoint. “All rows” still changes only two attention layers. Saved-O replacement uses FP64 recomputation followed by BF16 rounding, as in Table 10. I.3

T HE SOURCE OF THE SAVED - OUTPUT ERROR

The first-row violations point to the forward softmax, which forms the exponential argument r = fma(x, a, − round32 (ma)) ,

z = 2r ,

(20)

where m is the unscaled row maximum and a includes the attention scale and the base-two conversion. Because the fused multiply-add does not round xa separately, r need not vanish even when x = m, a hazard known from prior work (PyTorch contributors, 2024). If the value product then uses a BF16 cast of z while the denominator keeps its FP32 value, a one-key row outputs   roundBF16 (z) ob0 ≃ roundBF16 v0 (21) z rather than v0 , and this error enters a subtraction whose exact result is zero. Prescribed operands confirm this model. With all 64 query coordinates equal to c2k and all key coordinates equal to one, for c ∈ {3, 5, 7}, k = 0, . . . , 15 and sequence lengths 1 and 128, the dot products are known exactly; these 96 configurations share 48 constructions. The scalar model reproduces every native first-row output bitwise, including the six configurations with o0 ̸= v0 (Figure 8). Subtracting the row maximum with an explicit rounded operation before the base-two scaling, with the rest of the source and the compiler settings unchanged, gives FA3-SBS, the forward repair used throughout; it supplies the forward of GProj but no backward projection. It restores o0 = v0 and dQ = 0 in all 96 prescribed cases and, on the four checkpoint-8000 captures, reduces the absolute dQ error by 10.8–43.0× and the dK error by 12.0–58.6× (Figure 5). Using this forward in all 24 layers lowers the median model-gradient norm at checkpoint 8000 from 5353.856 to 18.892, with mean losses of 3.072690 and 3.073105. I.4

W HAT THE FORWARD REPAIR LEAVES BEHIND

Crossing saved states shows that the improvement comes from the saved output alone. For each capture we pair the FA3 and FA3-SBS saved outputs with their LSEs and pass all 32 states to three backward binaries: identical states give bitwise identical dQ, dK, dV in all three, and on the late captures changing only O reduces the query-gradient error by 10.8–43.0×, while changing only the LSE slightly worsens every case. The repaired backward is still sensitive to where the keys sit. Setting one query coordinate to zero and the matching coordinate of every key to a common BF16 value b leaves scores, output and LSE bitwise unchanged, and the exact gradient in that coordinate is zero. The native error in that coordinate still grows linearly with b, with FA3-SBS and even with an exact saved output; at b = 65536, FA3 and FA3-SBS give 665.15 and 665.54 (Appendix J.3, Figure 9). This is the translationinvariance failure that Section 3.4 and the −24 witness of Appendix A trace to the BF16 cast, and that no forward repair can remove. The failure is not specific to FA3. At fixed captured operands FA2 reproduces the FA3 errors, while cuDNN and PyTorch’s fused SDPA backends behave like FA3-SBS and retain a substantial backward error (Appendix J.4). 23

Preprint

FA3

FA3-SBS

PG19

(1) Full-sequence query-gradient error

(2) Full-sequence key-gradient error 10−1

‖dK ̂ − dKFP64‖F

10−1

‖dQ ̂ − dQFP64‖F

ProofPile

10−2

10−3 23.1B L5

23.1B L11

33.6B L5

10−2

10−3

33.6B L11

23.1B L5

10

5353.86

3

102 101

18.89 2.62

2.36

FA3 23.1B

FA3-SBS 23.1B

100 FA3 33.6B

33.6B L11

(4) Full-model forward also changes

(3) Full-model gradient norm 104

33.6B L5

Training tokens / layer

Loss change (FA3-SBS minus FA3)

Single-document full-model gradient norm

Training tokens / layer

23.1B L11

FA3-SBS 33.6B

1e−3 2 1 0 −1

23.1B

33.6B Training tokens

(1,2): Fixed Q/K/V/dO, all 4096 queries; 2 documents × 2 layers × 2 checkpoints (23.1B, 33.6B tokens). (3,4): 8 documents per checkpoint; all 24 layers use the selected native forward and backward. Bars: gradient-norm medians; loss-change means.

Figure 5: Native source intervention on the from-scratch FA3 run. (1)–(2) Eight fixed captures and incoming derivatives, evaluating all 4096 query rows against FP64 attention. (3)–(4) Whole-model runs on eight documents per checkpoint, with the FA3-SBS forward used in all 24 attention layers; unlike the fixed-forward controls in Table 11, this changes the forward computation.

J

N UMERICAL D IAGNOSIS : D ETAILS AND A DDITIONAL C ONTROLS

J.1

M EASUREMENT SETS AND REFERENCES

Measurement sets. The local analysis of Appendix I.1 samples query rows 0, 32, . . . , 4064 against all 4096 keys, on eight exploratory documents and on 127 confirmation documents, at layers 5 and 11 of checkpoints 5500 and 8000 (23.1B and 33.6B tokens). The fixed-forward interventions of Appendix I.2 use all rows of eight documents at checkpoints 6500 and 8000. The kernel comparisons use the eight complete captures of Appendix E, with all 4096 query rows. References. All references differentiate exact attention in FP64 at the represented BF16 operands. Saved outputs are recomputed in FP64 and rounded to BF16, and saved LSEs are rounded to FP32, matching the native formats. Figures 6 and 7 show the full distributions. J.2

P RESCRIBED FMA CASES

The 48 constructions of Appendix I.3 set all 64 coordinates equal, so every dot product is known exactly, and run at lengths 1 and 128 for 96 configurations. The incoming derivative is nonzero only at the first query, whose support is a single key at both lengths. The scalar model uses a libm fused multiply-add, an FP32 base-two exponential, BF16 rounding of the value-product operand and an FP32 denominator; it matches all 96 native first-row outputs bitwise, including the six violations. 24

Preprint

(2) Layer 5: δ-model relative residual

(1) Layer 5: Native dQ relative error

100

2 × 10−2

10−2 6 × 10

−3

4 × 10−3

(4) Layer 11: Native dQ relative error

4 × 10−2

Residual / native error

Ratio (not percent)

104 103 102 101

3 × 10

−2

2 × 10−2

10−2 6 × 10

−3

4 × 10−3 3 × 10−3

23.1B 33.6B Checkpoint (training tokens)

(3) Layer 5: Direction cosine: cos(Eδ, E)

1.00 0.99999

0.98

0.99998

0.96 0.94 0.92 0.90

(5) Layer 11: δ-model relative residual (6) Layer 11: Direction cosine: cos(Eδ, E)

Cosine (median annotated)

101

Cosine (median annotated)

3 × 10−2

Residual / native error

Ratio (not percent)

102

23.1B 33.6B Checkpoint (training tokens)

1.00 0.99998

0.98

0.99990

0.96 0.94 0.92 0.90 23.1B 33.6B Checkpoint (training tokens)

Confirmation set: 127 documents; each point aggregates 128 queries × 16 heads. Bars: medians.

Figure 6: The complete 127-document confirmation distributions. Each point pools 128 sampled query rows and 16 query heads for one document, against all 4096 keys; bars are medians. Columns show relative native dQ error, the unexplained fraction after the saved-output prediction, and the error cosine.

Saved O only (BF16 oracle)

dQ error reduction factor (native error / intervention error)

(1) Layer 5

Saved LSE only (FP32 oracle)

(2) Layer 11

103

102

101

100 23.1B 33.6B Checkpoint (training tokens)

23.1B 33.6B Checkpoint (training tokens)

Confirmation set: 127 documents; same external forward and dO. Bars: medians. Above 1 = lower error.

Figure 7: Saved-state interventions on the same 127-document confirmation set and sampled query rows. Each point is a paired absolute-dQ-error reduction, with one as the no-improvement reference. Saved O is recomputed in FP64 and rounded to BF16; LSE is recomputed in FP64 and rounded to FP32. Native forward outputs and incoming derivatives remain fixed.

The explicit rounded subtraction of FA3-SBS restores both identities in all 96 prescribed cases, and on the eight real captures its first-row outputs also equal their values. 25

Preprint

(1) Native saved-output inconsistency

(2) Native spurious query derivative

1e−3

1e−3

8

Multiplier 3 Multiplier 5 Multiplier 7

1.2

‖dQ0‖F (exact value = 0)

7

max|O0 − V0|

6 5 4 3 2 1 0

1.0 0.8 0.6 0.4 0.2 0.0

2

9

2

12

2

15

2

18

2

21

2

24

29

Raw QK dot product (before attention scaling)

212

215

218

221

224

Raw QK dot product (before attention scaling)

FMA/P-rounding scalar emulation matches native O exactly in 96/96 cases (including all 6 altered outputs). L = 1 and L = 128 first-row results coincide; 48 paired constructions shown. Agreement is not an instruction trace.

Figure 8: Prescribed one-key analysis. The 96 executions share 48 query/key parameter constructions across two sequence lengths. Exact mathematics requires the output/value identity and a zero query gradient. The scalar arithmetic hypothesis reproduces the observed output in every case, including the six violations. J.3

S AVED - STATE CROSSES AND THE REMAINING GAUGE DEFECT

Table 12 separates the two saved quantities on the four late captures (Appendix I.4): the FA3-SBS saved output alone reproduces the whole improvement, while its LSE alone slightly worsens every case, as factors below one indicate. Every identical supplied state gives bitwise identical gradients across the three backward binaries. Saved O

Saved LSE

dQ error reduction

dK error reduction

FA3-SBS FA3 FA3-SBS

FA3 FA3-SBS FA3-SBS

10.822–42.971 0.99676–0.99806 10.815–43.005

11.968–58.601 0.99636–0.99780 11.954–58.562

Table 12: Ranges of paired absolute-error reduction on the four checkpoint-8000 full-row captures. The improvement from the source intervention is recovered by its saved output with the FA3 LSE. The gauge construction uses sequence length 128, four query heads, two KV heads, dimension 64 and scale 1/8. Query coordinate zero is identically zero, the matching coordinate of every key is set to b, another coordinate carries a score gap of 0, 1 or 64, and values and upstream derivatives are one fixed random draw. Across three binaries, all 144 byte-level comparisons confirm that output and LSE are unchanged as b varies, and in all 126 nonzero-offset comparisons the erroneous coordinate divided by b is exactly the same vector: the error is linear in b. At gap 64 the slopes are 0.0101494 for FA3 and 0.0101553 for FA3-SBS, which give the reported 665.15 and 665.54 at b = 65536 (Figure 9). An exact saved output still leaves a linear error, consistent with the −24 witness of Equation (8) (Appendix A), which isolates the cast channel with exact saved output and reduction. J.4

BACKEND COMPARISON AT IDENTICAL CAPTURED OPERANDS

Table 13 compares attention backends on the eight captures at fixed Q, K, V, U , against the FP64 reference, using PyTorch 2.8.0, CUDA 12.8 and cuDNN 9.10 on one H200; the efficient SDPA path needs explicit KV-head expansion, and cuDNN cannot run the length-one cases. FA2 reproduces FA3: per-capture relative errors differ by less than 3 × 10−8 , and it shows the same six one-key violations. cuDNN and PyTorch’s flash and efficient SDPA backends behave like FA3-SBS in the forward, yet their median relative query-gradient error remains about 2.2, against 0.003–0.004 for the materialized math paths, and every fused backend returns a nonzero first-row 26

(1) Saved-state transfer 23.1B tokens

Gap 0

105

33.6B tokens

101

100

Gap 1 Gap 64

104

SBS O only

SBS LSE only

1e−2

1.4

103 102 101 100

1.2

1.0

10−1 10−2

FA3 state

1.6

‖dQ⋅, c‖2/|b|

102

(3) Residual persists across states

(2) Spurious coordinate gradient

‖dQ⋅, c‖2 (exact value: 0)

Query-gradient error reduction factor

Preprint

FA3, native O

20

SBS O + LSE

212

0.8 0 1 64 Key-coordinate contrast gap

224

Common key coordinate b FA3, native O

FA3, oracle O

FA3-SBS, native O

FA3-SBS, oracle O

(1): 8 fixed Q/K/V/dO captures, all 4096 queries, FP64 reference. Backward binaries agree at fixed O/LSE. (2,3): Q⋅, c = 0 and K⋅, c = b; native O and LSE are unchanged across b. The normalized coordinate gradient is exactly constant at all seven nonzero offsets. At b = 0 it is zero. Oracle O is the FP64 output cast to BF16. Only the selected coordinate has an exact-zero gradient.

Figure 9: Two controls. (1) Saved-output/LSE crosses on fixed real captures isolate the observed improvement to saved output. (2)–(3) A synthetic common-key coordinate changes no score or native forward state but amplifies a mathematically zero query-gradient coordinate. Backend FA3 FA3-SBS FA2, 2.8.3.post1 SDPA cuDNN SDPA flash SDPA efficient, KV expanded SDPA math, BF16 SDPA math, FP32

dQ error dK error dQ cosine Late L11 dQ error One-key violations 7.729 2.187 7.729 2.206 2.217 2.147 0.0037 0.0031

0.8555 0.1327 0.8555 0.1287 0.1289 0.1271 0.0035 0.0032

0.131 0.414 0.131 0.412 0.410 0.422 1.000 1.000

162.3 4.12 162.3 4.16 3.90 4.04 0.0077 0.0074

6/96 0/96 6/96 0/48 0/96 0/96 0/96 0/96

Table 13: Fixed-operand backend sweep on the from-scratch FA3 run’s captures. Errors are relative L2 ratios, not percentages; the first three numeric columns are medians over eight full-row captures. “Late L11” is the checkpoint-8000 PG19 layer-11 capture. One-key violations count prescribed cases with o0 ̸= v0 (the same cases have nonzero dq0 ); cuDNN cannot run the 48 length-one cases. dQ on all eight captures, including those with o0 = v0 . The backward error is thus shared across fused BF16 backends.

K

A P LAIN -T RANSFORMER C ONTROL

We train three independently initialized standard-residual transformers (102M parameters; 12 layers, width 512, eight query heads, two KV heads, head dimension 64, no QK normalization) for 8,000 updates at batch 8 and sequence length 4096, each at a base learning rate (LR1) and at a fourfold stress (LR4). Periodically, and at training-batch gradient peaks, we replace every attention backward by an FP32 reference while keeping all attention outputs and the loss bitwise fixed. In this model the attention backward stays close to the FP32 reference. At 500 updates the full-model gradient discrepancy is 0.26–0.27% at LR1 and 1.5–2.4% at LR4 (Table 14), and the worst values over all periodic checks are 0.88% and 12.4%. The LR4 runs do become unstable, with every final update clipped, training NLL of 5.34–5.40 against 3.54–3.57 at LR1, and gradient-norm peaks of 2,152–4,914. At the peak of seed 2026 (update 5054), the FP32 backward gives a norm of 4,988.69 against 4,913.88 natively (Table 15). The difference from the 450M FA3 run lies in scale: the plain models’ maximum query RMS stays at 42.5–44.5, against above 200 in layer 11 of the FA3 run at 33.6B tokens (Figure 10).

27

Preprint

Seed

LR

Rel. err. 0

Rel. err. 250

Rel. err. 500

Abs. err. 500

Cosine 500

2026 2026 2027 2027 2028 2028

1 4 1 4 1 4

0.00177832 0.00177832 0.00193911 0.00193911 0.00186521 0.00186521

0.00187079 0.00340782 0.00194124 0.00345874 0.00200166 0.00360284

0.00261074 0.0243191 0.00264846 0.0156458 0.0026835 0.015028

0.00498338 0.0514253 0.00479022 0.0314603 0.00528956 0.0348366

0.999996597 0.999706727 0.999996495 0.999878256 0.999996407 0.999897825

Table 14: Full-model gradient analyses of three independent 101,986,816-parameter plain transformer initializations. All attention forward outputs and loss are bitwise matched; only all-attention backward is replaced by FP32 reference arithmetic. Relative error uses the reference gradient norm. Seed Peak step Native norm Reference norm Rel. error 2026

5054

4913.875

4988.690

Cosine

0.018294 0.999944

Table 15: The LR4 training-batch peak of seed 2026. The analysis replaces every attention backward on the identical native forward and accumulates all eight sequence gradients in FP32 with training’s scored-token weights. Three initializations; same native forward; LR at 1% floor after dashed line

10−1

LR 1 LR 4 (preset stress)

10−2

50 100 150 200 250 Cumulative training token slots (millions)

(2) Naturally learned scale vs. error

Full-model relative gradient error

Full-model relative gradient error

(1) Gradient error over training

10−1

10−2

3

5 10 20 Maximum over layers of Q RMS

40

Figure 10: All six completed plain-transformer trajectories. The dashed line marks the learningrate-floor transition at update 4000.

28

Record · ID 1108716 · SHA-256 453071638b89b2d6
Retrieved via Conceptio — every document is proof-bundled with source, license, and retrieval metadata.