{"slug": "fa3-had-some-bugs", "title": "FA3 had some bugs", "summary": "Researchers found two bugs in FlashAttention 3 (FA3) that caused gradient norm to grow a thousandfold and loss to end 0.2 nats above FP32 attention when pretraining a 450M-parameter transformer on 50B tokens, with training healthy for the first 25B tokens and no NaN appearing. One bug, carried over from FA2, is an fma rounding issue in the exp(x_i * scale - max_scaled) computation, and the second is that FA3 rounds dS to BF16 before multiplying it, so the attention-score gradient no longer sums to exactly zero. The issues surfaced because the team was using Muon rather than AdamW, whose weight decay and common QK-norm or QK-clipping practices had masked the problem.", "body_md": "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.\n\nLean, the language the proofs are written in, originally started for proving software so you might think we were using it for that too?\n\nBut, 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:\n\nWhen 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.\n\nThere 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:\n\nWe’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.\n\nNow, for the largest value of `x_i`, `max_scaled = x_i * scale`. So, `x_i * scale - max_scaled == 0`, right?\n\nUnfortunately, 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!\n\nThis was an opt-in thing for FA2, and just never fixed for FA3.\n\nTurns 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.\n\nPart 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.\n\nQuasi-Riemann is one thing, but reduced precision floats? Quite another.", "url": "https://wpnews.pro/news/fa3-had-some-bugs", "canonical_source": "https://ianbarber.blog/2026/10/08/fa3-had-some-bugs/", "published_at": "2026-10-09 04:25:28+00:00", "updated_at": "2026-10-09 04:49:02.729022+00:00", "lang": "en", "topics": ["machine-learning", "ai-research", "large-language-models", "ai-infrastructure"], "entities": ["OpenAI", "FlashAttention 3", "FlashAttention 2", "PyTorch", "Muon", "AdamW", "Lean"], "also_reported_by": [], "alternates": {"html": "https://wpnews.pro/news/fa3-had-some-bugs", "markdown": "https://wpnews.pro/news/fa3-had-some-bugs.md", "text": "https://wpnews.pro/news/fa3-had-some-bugs.txt", "jsonld": "https://wpnews.pro/news/fa3-had-some-bugs.jsonld"}}