FA3 had some bugs

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 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. 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 round_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.

Discover more from Ian’s Blog

Subscribe now to keep reading and get access to the full archive.

Continue reading