Performance
Flash Attention : de FA1 à FA4
Flash Attention optimise l'attention en minimisant les transferts HBM/SRAM : tiling, recomputation, et l'évolution FA1 à FA4 (Ampere → Blackwell).
Flash Attention est un algorithme qui optimise le mécanisme d’attention en minimisant les transferts mémoire entre la HBM (RAM du GPU) et la SRAM (cache rapide), offrant jusqu’à 3x d’accélération sans aucune approximation.
Le problème de l’attention standard
Le mécanisme d’attention nécessite de calculer une matrice N×N (où N = longueur de séquence), ce qui pose deux problèmes majeurs :
- Complexité mémoire O(N²) : la matrice d’attention explose pour les longues séquences.
- Bande passante mémoire : les nombreux allers-retours HBM ↔ SRAM créent un goulot d’étranglement.
Architecture mémoire GPU
Pour comprendre Flash Attention, il faut connaître la hiérarchie mémoire des GPU :
- HBM (High Bandwidth Memory) : la « RAM » du GPU. Grande capacité (40-80 Go sur A100/H100) mais bande passante « limitée » (~1,5-2 To/s). C’est ici que résident les poids du modèle.
- SRAM (on-chip) : le « cache L1/L2 » du GPU. Très petite capacité (~64-228 Ko par SM selon le GPU) mais bande passante énorme (~19 To/s). C’est où les calculs doivent avoir lieu.
L’attention standard transfère les données HBM ↔ SRAM plusieurs fois : charger Q/K/V, écrire QKᵀ, recharger pour softmax, écrire les scores, recharger pour multiplier par V… Flash Attention fait tout en un seul passage.
Comment fonctionne Flash Attention
1. Tiling (découpage en blocs)
Au lieu de calculer l’attention sur toute la séquence d’un coup, Flash Attention découpe Q, K, V en petits blocs qui tiennent dans la SRAM. L’attention est calculée bloc par bloc, et les résultats sont combinés progressivement.
┌─────────────┐ ┌─────────────┐ ┌─────────────┐
│ Q (N×d) │ │ Kᵀ (d×N) │ │ V (N×d) │
├──┬──┬──┬──┬─┤ ├──┬──┬──┬──┬─┤ ├──┬──┬──┬──┬─┤
│Q₁│Q₂│Q₃│Q₄│…│ x │K₁│K₂│K₃│K₄│…│ x │V₁│V₂│V₃│V₄│…│
└──┴──┴──┴──┴─┘ └──┴──┴──┴──┴─┘ └──┴──┴──┴──┴─┘
↓ ↓ ↓
Blocs chargés séquentiellement en SRAM
2. Recomputation (recalcul intelligent)
Au lieu de stocker la matrice d’attention complète (O(N²) en mémoire), Flash Attention stocke uniquement les facteurs de normalisation softmax (O(N)). Lors du backward pass, l’attention est recalculée à la volée depuis ces facteurs.
Résultats de performance
| Métrique | Attention standard | Flash Attention | Gain |
|---|---|---|---|
| Complexité mémoire | O(N²) | O(N) | Linéaire |
| Vitesse (seq 2K) | 1x | 2-3x | +200-300 % |
| Vitesse (seq 8K) | 1x | 2,7x | +270 % |
| Longueur max supportée | ~4K-8K | 64K+ | ×8-16 |
Flash Attention 2, 3 et 4
Les versions successives apportent des optimisations supplémentaires, en grande partie liées à l’évolution des architectures GPU NVIDIA (Ampere → Hopper → Blackwell).
| Version | Année | GPU cible | Précision | Gain vs précédent |
|---|---|---|---|---|
| FA1 | 2022 | A100 | FP16/BF16 | ~3x vs attention naïve |
| FA2 | 2023 | A100/H100 | FP16/BF16 | ~2x vs FA1 (parallélisme amélioré) |
| FA3 | 2024 | H100 (Hopper) | FP16/BF16/FP8 | ~1,5-2x vs FA2 (TMA + warp-specialization) |
| FA4 | 2025 | B200 (Blackwell) | FP8/FP4 | ~1,6x vs FA3 (5e gen Tensor Cores) |
Flash Attention 2 (2023)
FA2 améliore radicalement la stratégie de parallélisation : la boucle externe parcourt les tuiles de Q (et non plus de K), permettant d’attribuer une tuile à chaque thread block. Sur H100, FA2 atteint ~70 % du débit FP16 théorique du GPU.
Flash Attention 3 (2024)
FA3 exploite les nouveautés de l’architecture Hopper :
- TMA (Tensor Memory Accelerator) : transferts asynchrones HBM↔SMEM en parallèle des calculs.
- Warp-specialization : certains warps font le chargement, d’autres le calcul, en pipeline.
- FP8 natif : ×2 sur le débit grâce aux Tensor Cores FP8 (E4M3/E5M2).
- 740 TFLOPS en FP16 et 1,2 PFLOPS en FP8 sur H100, soit 75 % du peak théorique.
Flash Attention 4 / dFlash (2025-2026)
Annoncé en 2025 et déployé courant 2026, FA4 cible les GPU Blackwell (B100/B200, GB200) et leur 5ᵉ génération de Tensor Cores avec support natif FP4 (NVFP4/MXFP4). La variante dFlash (« dynamic Flash ») ajoute des optimisations spécifiques à l’inférence à long contexte :
- FP4 attention : softmax et multiplication des scores en 4 bits, 2x plus de débit qu’en FP8.
- Block-sparse attention dynamique : le pattern de sparsité est calculé à la volée selon les scores softmax — les blocs inutiles ne sont jamais matérialisés (gain ×2-×4 sur les contextes > 64K).
- Online softmax fusionné avec rotary embeddings : un seul kernel pour Q·Kᵀ, RoPE et softmax.
- Decode-optimized variants : kernels dédiés à la phase de génération token-par-token, où Q est de longueur 1 (FA4-decode) — gain ×3 sur le decoding pur.
- Support natif des architectures MLA (DeepSeek) et NSA (Native Sparse Attention).
En parallèle de FA4, des bibliothèques comme FlashInfer (CMU) et SageAttention proposent des kernels d’attention quantifiée (INT8/INT4) pour l’inférence. SageAttention 2 atteint des accélérations 2-3x sur RTX 4090 avec une perte de qualité négligeable, ce qui en fait une alternative populaire pour le serving local.
Sliding Window et Native Sparse Attention
Les modèles 2026 (Qwen3.6, Mistral Small 3.2, Gemma 3) combinent souvent Flash Attention avec des patterns d’attention sparse :
- Sliding Window : chaque token n’attend que les W derniers — utilisé dans Mistral, en alternance avec attention complète tous les N layers.
- NSA (Native Sparse Attention) : DeepSeek V3.2, gain ×10 sur des contextes > 128K avec impact qualité < 1 %.
- Differential Attention (Microsoft, 2024) : deux têtes en différentiel pour réduire le bruit d’attention sur les longs contextes.
Utilisation pratique
# Activer Flash Attention dans llama.cpp
./llama-cli -m model.gguf --flash-attn
# Compiler avec support Flash Attention (nécessaire pour cache V quantifié)
cmake -B build -DGGML_CUDA_FA_ALL_QUANTS=ON -DGGML_CUDA=ON
cmake --build build
from transformers import AutoModelForCausalLM
import torch
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3.6-27B-Instruct",
torch_dtype=torch.bfloat16,
attn_implementation="flash_attention_3" # FA3 si SM90+ (H100/B200)
)
Flash Attention et le cache KV sont complémentaires : le cache KV évite les recalculs redondants, tandis que Flash Attention optimise le calcul d’attention restant. C’est aussi Flash Attention qui rend possible la quantification du cache V dans llama.cpp (--flash-attn est requis pour --cache-type-v). Utilisés ensemble, ils permettent des inférences extrêmement rapides même sur de très longues séquences.