Ce que FlashAttention change dans l’inférence des transformers
Traduction automatique
Cet article a été traduit automatiquement depuis la version originale en anglais.
FlashAttention calcule l’attention exacte sans stocker les matrices complètes de scores et de probabilités d’attention dans la mémoire principale du GPU. Il traite des blocs plus petits et combine leurs résultats avec un softmax en ligne. Cela réduit les transferts mémoire et le stockage intermédiaire, tout en conservant le coût arithmétique quadratique de l’attention dense en fonction de la longueur de séquence.
Une moindre utilisation de mémoire ne rend pas linéaire la croissance du coût de calcul de l’attention dense.
Distinguer stockage et calcul
L’attention matérialisée standard calcule des scores pour chaque paire requête/clé visible, normalise ces scores et les multiplie par les valeurs. Une matrice d’attention carrée sans masque contient N² entrées pour N tokens.
L’article original de FlashAttention réorganise ce calcul en blocs qui tiennent sur la puce. Il conserve des statistiques de normalisation et des accumulateurs de sortie mis à jour au fil du calcul, au lieu d’écrire toutes les probabilités en HBM. À dimensions de tête fixes, le stockage nécessaire aux résultats intermédiaires d’attention croît linéairement avec la longueur de séquence. Il évalue toujours les interactions denses entre requêtes et clés qu’exige le masque d’attention.
L’algorithme calcule une attention exacte ; des ordres d’exécution différents des opérations en virgule flottante peuvent entraîner de faibles différences numériques. Ce n’est pas une méthode approximative d’attention parcimonieuse qui écarte certaines paires de tokens.
L’implémentation et la charge de travail comptent toujours
| Question | Pourquoi la vérifier |
|---|---|
| Le backend prend-il en charge ce GPU et cette précision ? | Les implémentations des kernels ciblent des matériels précis |
| La dimension des têtes, le masque et l’organisation du cache sont-ils pris en charge ? | Les formes non prises en charge peuvent entraîner la sélection d’un autre backend |
| L’attention représente-t-elle une grande part du temps de requête ? | Un gain important au niveau du kernel peut avoir un faible effet global |
| La charge correspond-elle au prefill ou au décodage d’un seul token ? | Leur travail parallélisable diffère |
FlashAttention-2 parallélise aussi les blocs de la séquence de requêtes, en plus des lots et des têtes, et réduit d’autres surcoûts d’exécution. FlashAttention-3 cible des fonctions d’exécution propres à Hopper. Les noms de version décrivent des algorithmes et des implémentations ; ils ne permettent pas de déterminer le backend réellement sélectionné par votre serveur.
Vérifiez le backend sélectionné dans le runtime de serving dont la version est fixée. Mesurez la latence du prompt, le pic de mémoire et le débit de bout en bout avec des longueurs et des tailles de lot représentatives. Comparez des masques et une précision identiques. Un benchmark qui ne rapporte que les TFLOPS du kernel ne démontre pas une amélioration de la latence par token perçue par l’utilisateur.
Guide d’ingénierie : FlashAttention présente les versions successives et leurs mesures dans leurs périmètres respectifs.