3 · FlashAttention¶
L'algorithme qui a rendu le contexte long possible, et le meilleur exemple existant de co-conception algorithme/matériel.
3.1 Le problème¶
L'attention standard :
où \(\mathbf{Q}, \mathbf{K}, \mathbf{V} \in \mathbb{R}^{S \times d_h}\) et \(\mathbf{S}, \mathbf{P} \in \mathbb{R}^{S \times S}\).
Le problème n'est pas le nombre d'opérations, c'est la matrice intermédiaire.
Pour \(S = 8192\) en FP16 : \(\mathbf{S}\) occupe \(8192^2 \times 2 = 134\) Mo. Par tête. Avec 32 têtes et un lot de 8 : 34 Go, écrits puis relus deux fois (pour le softmax, puis pour la multiplication par \(\mathbf{V}\)).
Décompte du trafic pour l'attention standard, par tête :
| Étape | Octets |
|---|---|
| Lire \(\mathbf{Q}\), \(\mathbf{K}\) | \(4 S d_h\) |
| Écrire \(\mathbf{S}\) | \(2 S^2\) |
| Lire \(\mathbf{S}\), écrire \(\mathbf{P}\) | \(4 S^2\) |
| Lire \(\mathbf{P}\), \(\mathbf{V}\) | \(2S^2 + 2Sd_h\) |
| Écrire \(\mathbf{O}\) | \(2 S d_h\) |
| Total | \(\approx 8S^2 + 8Sd_h\) |
Pour \(S = 8192\) et \(d_h = 128\) : \(5{,}4 \times 10^8\) octets contre \(3{,}4\times10^{10}\) FLOP, soit \(I \approx 63\). En dessous du seuil de 296.
3.2 L'idée : ne jamais matérialiser S¶
FlashAttention (Dao et al., 2022) calcule l'attention par blocs, en gardant tout en mémoire partagée, sans jamais écrire \(\mathbf{S}\) en HBM.
Pour chaque bloc de lignes de Q (taille Br) :
charger Q_i en mémoire partagée
initialiser O_i = 0, l_i = 0, m_i = -∞
Pour chaque bloc de colonnes de K, V (taille Bc) :
charger K_j, V_j en mémoire partagée
S_ij = Q_i · K_jᵀ ← en registres/mémoire partagée
mettre à jour m_i, l_i, O_i ← softmax EN LIGNE
écrire O_i en HBM
Le trafic devient :
Le terme quadratique disparaît : on ne relit que \(\mathbf{K}\) et \(\mathbf{V}\), et seulement \(S/B_r\) fois.
Avec \(S = 8192\), \(d_h = 128\), \(B_r = 128\) : trafic \(\approx 2{,}7\times10^8\) octets, soit 2 fois moins que l'attention standard — et surtout, la croissance devient linéaire en \(S\) au lieu de quadratique.
3.3 Le softmax en ligne¶
C'est la clé algorithmique. Comment normaliser sans avoir vu toutes les valeurs ?
Le softmax stable exige \(\max_j x_j\) et \(\sum_j \exp(x_j - \max)\), qui ne sont connus qu'après avoir tout parcouru.
La solution : maintenir des statistiques courantes et corriger rétroactivement.
Après avoir traité les blocs \(1 \dots j\), on maintient :
Le facteur \(e^{m^{(j-1)} - m^{(j)}}\) remet à l'échelle l'accumulation précédente quand un nouveau maximum apparaît. À la fin, on divise par \(\ell^{(S/B_c)}\).
Pourquoi c'est correct
Si le maximum passe de \(m\) à \(m'\), chaque terme précédemment calculé comme \(e^{x - m}\) aurait dû être \(e^{x - m'} = e^{x-m} \cdot e^{m - m'}\).
Multiplier l'accumulation par \(e^{m-m'}\) corrige donc exactement tous les termes d'un coup. Le résultat est numériquement identique au softmax global, pas une approximation.
C'est une propriété rare et précieuse : l'algorithme est exact, pas approché.
3.4 Le recalcul en rétropropagation¶
Pour la passe arrière, il faut \(\mathbf{P}\). Or on ne l'a pas stockée.
FlashAttention la recalcule à partir de \(\mathbf{Q}\), \(\mathbf{K}\) et des statistiques \((m, \ell)\) sauvegardées — qui ne font que \(O(S)\) au lieu de \(O(S^2)\).
C'est un arbitrage explicite : plus de FLOP contre moins d'octets. Puisque l'attention est limitée par la mémoire, c'est gagnant.
Recalculer \(\mathbf{S}\) coûte \(2S^2 d_h\) FLOP supplémentaires ; le stocker coûterait \(4S^2\) octets d'écriture + lecture. Le rapport penche largement du côté du recalcul.
Une leçon générale
« Recalculer plutôt que stocker » est un principe qui dépasse FlashAttention : c'est aussi celui du gradient checkpointing, du pavage temporel des stencils, et de nombreuses optimisations de mémoire.
La règle : quand \(I \ll I_{\text{crit}}\), échangez toujours des FLOP contre des octets.
3.5 Les versions¶
| Version | Année | Apport principal | Utilisation matérielle |
|---|---|---|---|
| FA1 | 2022 | pavage + softmax en ligne + recalcul | ~25-40 % |
| FA2 | 2023 | meilleure répartition du travail, moins d'opérations non-matmul, parallélisation sur la longueur de séquence | ~50-73 % |
| FA3 | 2024 | TMA, wgmma, chevauchement softmax/GEMM, FP8 |
~75 % (H100) |
| FA4 | 2026 | co-conception pour Blackwell | 71 % (B200) |
FlashAttention-2¶
Le constat : dans FA1, les opérations non-matmul (exponentielles, remises à l'échelle, divisions) consomment une part disproportionnée du temps, parce qu'elles s'exécutent sur les unités vectorielles, ~16× plus lentes que les tensor cores.
Les corrections : réduire le nombre de remises à l'échelle, paralléliser sur la dimension de séquence (et pas seulement sur le lot et les têtes), et mieux répartir le travail entre warps.
FlashAttention-3¶
Exploite Hopper : TMA pour les chargements, wgmma asynchrone, et surtout un
chevauchement explicite entre la GEMM du bloc \(j+1\) et le softmax du bloc
\(j\) — puisque wgmma est asynchrone, on peut lancer la première pendant qu'on
calcule le second sur les unités vectorielles.
FlashAttention-4¶
Publié le 5 mars 2026, avec des résultats préliminaires présentés à Hot Chips en août 2025.
Le constat de départ : sur Blackwell, les unités ne progressent pas au même rythme. Le débit des tensor cores double, tandis que la bande passante de la mémoire partagée et le débit des unités exponentielles (SFU) stagnent. L'attention, qui alterne GEMM et exponentielles, devient limitée par la SFU.
Les trois réponses :
- Exponentielles émulées en logiciel.
exp()est calculée par approximation polynomiale sur les unités FMA au lieu de la SFU. Blackwell a des FMA en abondance ; la SFU non. - Remise à l'échelle conditionnelle. Le softmax en ligne remet normalement à l'échelle chaque fois que le maximum courant change. FA4 saute cette remise à l'échelle tant que le décalage ne menace pas la stabilité numérique, ce qui réduit le nombre de rescalings d'un facteur ~10×.
- Implémentation entièrement en CuTe DSL, avec une compilation qui prend des secondes au lieu de minutes ou d'heures.
Résultats : jusqu'à 1 605 TFLOPS sur B200 (71 % d'utilisation matérielle), 1,3× plus rapide que cuDNN 9.13 et 2,7× plus rapide que les implémentations Triton.
Pourquoi FA4 est un cas d'école
L'algorithme mathématique n'a pas changé depuis 2022. Ce qui a changé, c'est l'appariement entre les opérations et les unités matérielles disponibles.
Déplacer exp() de la SFU vers les FMA n'est pas une optimisation
algorithmique : c'est une observation sur le rapport de débits entre unités
d'une puce précise. C'est exactement ce que signifie « co-conception ».
3.6 L'attention en décodage¶
Le décodage a un profil complètement différent : \(\mathbf{Q}\) est un seul jeton (ou \(b\) jetons), alors que \(\mathbf{K}\) et \(\mathbf{V}\) font toute la longueur du contexte.
Ce n'est plus une GEMM mais un produit matrice-vecteur : \(I \approx 1\), massivement limité par la mémoire, et les tensor cores sont inutilisables tels quels.
Flash-Decoding répond au problème du parallélisme insuffisant : avec un seul jeton de requête, il n'y a qu'un bloc de travail, ce qui laisse le GPU vide. La solution est de découper le contexte en morceaux traités en parallèle, puis de combiner les résultats partiels avec la même formule de fusion de softmax que FlashAttention.
C'est ce qu'implémentent FlashInfer et les noyaux d'attention de vLLM et SGLang, avec des variantes selon la disposition du cache (paged attention).
3.7 Les variantes à connaître¶
| Variante | Idée | Usage |
|---|---|---|
| Paged Attention | cache KV en pages non contiguës | vLLM, évite la fragmentation |
| Flash-Decoding | découpage du contexte pour le parallélisme | décodage à petit lot |
| FlexAttention (PyTorch) | fonction de masque/score arbitraire, compilée | masques exotiques |
| Ring Attention | répartition du contexte sur plusieurs GPU | contexte très long |
| Attention creuse | ne calculer que certains blocs | contexte très long |
| MLA (DeepSeek) | compression du cache par projection latente | réduction du cache KV |
Avant d'écrire un noyau d'attention
Vérifiez FlexAttention et FlashInfer. Les deux prennent des fonctions de score et de masque arbitraires et génèrent le noyau. Le nombre de cas justifiant un noyau manuel est bien plus faible qu'on ne le croit.
Résumé du chapitre¶
À retenir
- Le problème de l'attention n'est pas le calcul mais la matrice intermédiaire \(\mathbf{S}\) de taille \(S^2\).
- FlashAttention la ne matérialise jamais : pavage + softmax en ligne. Le trafic devient linéaire en \(S\).
- Le softmax en ligne est exact, pas approché : le facteur \(e^{m-m'}\) corrige rétroactivement toute l'accumulation.
- En rétropropagation, on recalcule \(\mathbf{P}\) plutôt que de la stocker. Principe général : quand \(I \ll I_{\text{crit}}\), échangez des FLOP contre des octets.
- FA4 (mars 2026) : exponentielles sur les unités FMA au lieu de la SFU, remise à l'échelle conditionnelle (~10× moins), écrit en CuTe DSL. 1 605 TFLOPS sur B200, 1,3× cuDNN, 2,7× Triton.
- Le décodage est un autre problème : produit matrice-vecteur, \(I \approx 1\), résolu par Flash-Decoding et les noyaux de FlashInfer.
Vérifiez que vous avez compris¶
Pourquoi le softmax en ligne donne-t-il un résultat exact et pas approché ?
Parce que la correction est algébriquement exacte.
Supposons qu'on ait accumulé \(\sum_i e^{x_i - m}\) avec l'ancien maximum \(m\), et qu'un nouveau maximum \(m' > m\) apparaisse. La valeur correcte serait \(\sum_i e^{x_i - m'}\).
Or :
Un seul facteur multiplicatif corrige tous les termes simultanément. Il n'y a aucune approximation, seulement de l'arithmétique flottante ordinaire.
(FA4 introduit un choix : sauter la remise à l'échelle quand \(m - m'\) est assez petit pour ne pas menacer la stabilité. Là, il y a bien un compromis, mais borné et contrôlé.)
Pourquoi FlashAttention est-il particulièrement efficace avec un masque causal ?
Parce qu'il peut sauter des blocs entiers.
Avec un masque causal, un bloc de requêtes \(i\) n'a besoin que des blocs de clés \(j \le i\). La boucle interne devient :
pour j de 0 à i : ← au lieu de 0 à S/Bc
On économise exactement la moitié des blocs — facteur 2 sur le calcul et le trafic, gratuitement.
L'attention standard, elle, calcule \(\mathbf{S}\) entièrement puis applique le masque : elle paie tout et en jette la moitié. C'est un exemple direct du point 1.5 de la checklist : ne pas calculer ce qui est masqué.
Sur Blackwell, FA4 calcule exp() sur les unités FMA plutôt que sur la SFU. Cela coûte plus d'instructions. Pourquoi est-ce plus rapide ?
Parce que le goulot n'est pas le nombre d'instructions mais le débit de l'unité saturée.
Blackwell a environ deux fois plus de débit tensor core que Hopper, mais un nombre d'unités SFU inchangé. Dans un noyau d'attention qui alterne GEMM (tensor cores) et exponentielles (SFU), la SFU devient le facteur limitant : les tensor cores attendent.
Déplacer les exponentielles vers les unités FMA — nombreuses, sous-utilisées dans ce noyau — coûte davantage d'instructions sur une ressource qui n'est pas saturée. Le temps total baisse.
C'est un raisonnement de roofline appliqué au niveau de l'unité fonctionnelle plutôt qu'au niveau du GPU entier, et c'est exactement l'un des cas où le modèle roofline simple ne suffit pas (voir Performance 1, §1.5).
Chapitre suivant : 4 · La fusion de noyaux
Sources de ce chapitre¶
- Dao, Fu, Ermon, Rudra, Ré, FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS 2022 — arXiv:2205.14135
- Dao, FlashAttention-2 — arXiv:2307.08691
- Shah et al., FlashAttention-3 — arXiv:2407.08608
- FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling — arXiv:2603.05451
- We reverse-engineered Flash Attention 4, Modal
- Flash-Decoding for long-context inference, PyTorch blog
- FlashInfer · FlexAttention, PyTorch blog