# FA3 had some bugs

> Source: <https://ianbarber.blog/2026/10/08/fa3-had-some-bugs/>
> Published: 2026-10-09 04:25:28+00:00

For reasons best known to themselves OpenAI blessed the world with several hundred new mathematical proofs. This inspired elation of some mathematicians, despair of many others, and confusion of a great many people whose only encounter with zeta moments involved Sean Connery and a laser grid.

Lean, the language the proofs are written in, originally started for proving software so you might think we were using it for that too?

But, surprisingly, [we are still finding bugs](https://arxiv.org/abs/2609.34272) in possibly the most widely used kernel of the last few years, FlashAttention 3:

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.

There were two issues. One was an issue I remember from [2024 in FA2](https://github.com/pytorch/pytorch/issues/121558). Horace’s comment on that issue sums it up very nicely:

We’re computing `exp(x_i * scale - max_scaled)`. Now, `max_scaled = max(x_i * scale)`. The idea here is that before we take the exponent, we normalize everything down to 0 so that `exp(x)` doesn’t blow up massively.

Now, for the largest value of `x_i`, `max_scaled = x_i * scale`. So, `x_i * scale - max_scaled == 0`, right?

Unfortunately, no, due to fma. `max_scaled` is actually equal to r`ound_fp32(x_i * scale)`. But in fma, `x_i * scale` never gets rounded! So with fma, we are computing `x_i * scale - round_fp32(x_i * scale)`, which is not equal to 0!

This was an opt-in thing for FA2, and just never fixed for FA3.

Turns out though, there was another bug: `dS`, which is the gradient with respect to the attention scores, should sum to exactly zero: its expresses how the keys differ from each other, not where they are absolutely. FA3 rounds `dS` to BF16 before multiplying it, so it doesn’t *quite* sum to zero. Early in training that’s mostly noise, but as the keys get larger and attention spikier it can swamp the real gradient.

Part of the reason this hadn’t come up was that applying QK-norm (or doing QK-clipping) is quite common, and AdamW’s weight decay tended to avoid the explosion too. Muon, however, is more exposed, and that’s what the team were using when they noticed this.

Quasi-Riemann is one thing, but reduced precision floats? Quite another.
