Modèles et plateformes d’IA
Flash Attention : Révolutionner l’efficacité des modèles de transformation

Par
Aayush Mittal Mittal
Alors que les modèles de transformation grandissent en taille et en complexité, ils rencontrent des défis importants en termes d’efficacité computationnelle et d’utilisation de la mémoire, en particulier lorsqu’ils traitent des séquences longues. Flash Attention est une technique d’optimisation qui promet de révolutionner la façon dont nous mettons en œuvre et mettons à l’échelle les mécanismes d’attention dans les modèles de transformation.
Dans ce guide complet, nous allons plonger dans les profondeurs de Flash Attention, en explorant ses concepts fondamentaux, les détails de sa mise en œuvre et l’impact profond qu’il a sur le domaine de l’apprentissage automatique.
Le problème : l’attention est coûteuse
Avant de plonger dans la solution, comprenons d’abord le problème que Flash Attention cherche à résoudre. Le mécanisme d’attention, bien que puissant, comporte un coût computationnel important, en particulier pour les séquences longues.
L’attention standard : un rappel rapide
Le mécanisme d’attention standard dans les modèles de transformation peut être résumé par l’équation suivante :
Attention(Q, K, V) = softmax(QK^T / √d) VOù Q, K et V sont respectivement les matrices de requête, de clé et de valeur, et d est la dimension des vecteurs de clé.
Bien que cette formulation soit élégante, sa mise en œuvre conduit à plusieurs incohérences :
- Goulots d’étranglement de la mémoire : La matrice d’attention intermédiaire (QK^T) a une taille de N x N, où N est la longueur de la séquence. Pour les séquences longues, cela peut rapidement épuiser la mémoire GPU disponible.
- Accès à la mémoire redondant : Dans les mises en œuvre standard, la matrice d’attention est calculée, stockée dans la mémoire à haute bande passante (HBM) et puis lue à nouveau pour l’opération softmax. Cet accès à la mémoire redondant est un goulet d’étranglement majeur.
- Sous-utilisation de la puissance de calcul de la GPU : Les GPU modernes ont une puissance de calcul (FLOPS) nettement supérieure à la bande passante de la mémoire. La mise en œuvre standard de l’attention est limitée par la mémoire, laissant une grande partie de la puissance de calcul de la GPU inutilisée.
Illustrons cela avec un exemple de code Python simple qui montre la mise en œuvre standard de l’attention :
&amp;lt;/pre&amp;gt; import torch <p>def standard_attention(Q, K, V): # Q, K, V shape: (batch_size, seq_len, d_model) d_k = K.size(-1) scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k)) attention_weights = torch.softmax(scores, dim=-1) return torch.matmul(attention_weights, V)</p>
Cette mise en œuvre, bien que simple, souffre des incohérences mentionnées ci-dessus. Le tenseur scores, qui a une forme (batch_size, seq_len, seq_len), peut devenir prohibitivement grand pour les séquences longues.
Entrer dans Flash Attention
Flash Attention, introduit par Tri Dao et ses collègues dans leur article de 2022, est une approche de calcul de l’attention qui réduit considérablement l’utilisation de la mémoire et améliore l’efficacité computationnelle. Les idées clés derrière Flash Attention sont :
- Tuiles : Diviser la grande matrice d’attention en petites tuiles qui tiennent dans la mémoire SRAM rapide.
- Recomputation : Au lieu de stocker la matrice d’attention entière, recomputer certaines parties pendant la passe arrière.
- Mise en œuvre consciente des entrées-sorties : Optimiser l’algorithme pour minimiser les mouvements de données entre les différents niveaux de la hiérarchie de mémoire de la GPU.
L’algorithme Flash Attention
Au cœur de Flash Attention, on retrouve une nouvelle façon de calculer le mécanisme d’attention. Au lieu de calculer la matrice d’attention entière à la fois, il la traite par blocs, en exploitant la hiérarchie de mémoire des GPU modernes.
Voici une vue d’ensemble de l’algorithme :
- Entrée : Matrices Q, K, V en HBM (High Bandwidth Memory) et dans la mémoire SRAM de taille M.
- Les tailles de bloc sont calculées en fonction de la mémoire SRAM disponible.
- Initialisation de la matrice de sortie O et des vecteurs auxiliaires l et m.
- L’algorithme divise les matrices d’entrée en blocs pour les faire tenir dans la mémoire SRAM.
- Deux boucles imbriquées traitent ces blocs :
- Boucle externe charge les blocs K et V
- Boucle interne charge les blocs Q et effectue les calculs
- Les calculs sur la mémoire SRAM incluent la multiplication matricielle, softmax et le calcul de la sortie.
- Les résultats sont écrits dans la HBM après le traitement de chaque bloc.
Ce calcul par bloc permet à Flash Attention de maintenir une empreinte mémoire nettement plus petite tout en calculant une attention exacte.
Les mathématiques derrière Flash Attention
La clé pour faire fonctionner Flash Attention est un astuce mathématique qui permet de calculer softmax de manière bloquée. L’article introduit deux formules clés :
- Décomposition softmax :
softmax(x) = exp(x - m) / Σexp(x - m)où m est la valeur maximale dans x.
- Fusion softmax :
softmax(x ∪ y) = softmax(softmax(x) * e^(m_x - m), softmax(y) * e^(m_y - m))où m = max(m_x, m_y)
Ces formules permettent à Flash Attention de calculer des résultats partiels de softmax pour chaque bloc et de les combiner correctement pour obtenir le résultat final.
Détails de la mise en œuvre
Plongeons dans une mise en œuvre simplifiée de Flash Attention pour illustrer ses concepts fondamentaux :
import torch <p>def flash_attention(Q, K, V, block_size=256): batch_size, seq_len, d_model = Q.shape</p> <p># Initialisation de la sortie et des statistiques en cours O = torch.zeros_like(Q) L = torch.zeros((batch_size, seq_len, 1)) M = torch.full((batch_size, seq_len, 1), float('-inf'))</p> <p>for i in range(0, seq_len, block_size): Q_block = Q[:, i:i+block_size, :]</p> <p>for j in range(0, seq_len, block_size): K_block = K[:, j:j+block_size, :] V_block = V[:, j:j+block_size, :]</p> <p># Calcul des scores d'attention pour ce bloc S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p> <p># Mise à jour de la valeur maximale en cours M_new = torch.maximum(M[:, i:i+block_size], S_block.max(dim=-1, keepdim=True).values)</p> <p># Calcul des exponentielles exp_S = torch.exp(S_block - M_new) exp_M_diff = torch.exp(M[:, i:i+block_size] - M_new)</p> <p># Mise à jour de la somme en cours L_new = exp_M_diff * L[:, i:i+block_size] + exp_S.sum(dim=-1, keepdim=True)</p> <p># Calcul de la sortie pour ce bloc O[:, i:i+block_size] = ( exp_M_diff * O[:, i:i+block_size] + torch.matmul(exp_S, V_block) ) / L_new</p> <p># Mise à jour des statistiques en cours L[:, i:i+block_size] = L_new M[:, i:i+block_size] = M_new</p> return O
Cette mise en œuvre, bien que simplifiée, capture l’essence de Flash Attention. Elle traite les entrées par blocs, en maintenant des statistiques en cours (M et L) pour calculer correctement softmax sur tous les blocs.
L’impact de Flash Attention
L’introduction de Flash Attention a eu un impact profond sur le domaine de l’apprentissage automatique, en particulier pour les grands modèles de langage et les applications à long contexte. Certains des avantages clés incluent :
- Utilisation réduite de la mémoire : Flash Attention réduit la complexité de la mémoire de O(N^2) à O(N), où N est la longueur de la séquence. Cela permet de traiter des séquences beaucoup plus longues avec le même matériel.
- Amélioration de la vitesse : En minimisant les mouvements de données et en utilisant mieux les capacités de calcul de la GPU, Flash Attention atteint des accélérations importantes. Les auteurs rapportent jusqu’à 3 fois plus de rapidité pour l’entraînement de GPT-2 par rapport aux mises en œuvre standard.
- Calcul exact : Contrairement à certaines autres techniques d’optimisation de l’attention, Flash Attention calcule une attention exacte, et non une approximation.
- Évolutivité : La réduction de l’empreinte mémoire permet de passer à des séquences beaucoup plus longues, potentiellement jusqu’à des millions de jetons.
Impact réel
L’impact de Flash Attention s’étend au-delà de la recherche académique. Il a été rapidement adopté dans de nombreuses bibliothèques et modèles d’apprentissage automatique populaires :
- Transformateurs Hugging Face : La bibliothèque de transformateurs populaire a intégré Flash Attention, permettant aux utilisateurs de profiter facilement de ses avantages.
- GPT-4 et au-delà : Même si cela n’est pas confirmé, il y a des spéculations selon lesquelles des modèles de langage avancés comme GPT-4 pourraient utiliser des techniques similaires à Flash Attention pour gérer les contextes longs.
- Modèles à long contexte : Flash Attention a permis l’émergence d’une nouvelle génération de modèles capables de gérer des contextes extrêmement longs, tels que des modèles qui peuvent traiter des livres entiers ou des vidéos longues.
FlashAttention : Développements récents
FlashAttention-2
En s’appuyant sur le succès de la version originale de Flash Attention, la même équipe a introduit FlashAttention-2 en 2023. Cette version mise à jour apporte plusieurs améliorations :
- Optimisation supplémentaire : FlashAttention-2 atteint une meilleure utilisation de la GPU, atteignant jusqu’à 70 % du pic théorique de FLOPS sur les GPU A100.
- Amélioration de la passe arrière : La passe arrière est optimisée pour être presque aussi rapide que la passe avant, conduisant à des accélérations importantes lors de l’entraînement.
- Prise en charge de différentes variantes d’attention : FlashAttention-2 étend la prise en charge à diverses variantes d’attention, y compris l’attention à requête groupée et l’attention à plusieurs requêtes.
FlashAttention-3
Publié en 2024, FlashAttention-3 représente la dernière avancée dans cette lignée de recherche. Il introduit plusieurs nouvelles techniques pour améliorer encore les performances :
- Calcul asynchrone : En exploitant la nature asynchrone des nouvelles instructions GPU pour chevaucher différents calculs.
- Prise en charge du FP8 : En utilisant le calcul à basse précision FP8 pour un traitement encore plus rapide.
- Traitement incohérent : Une technique pour réduire l’erreur de quantification lors de l’utilisation de formats à basse précision.
Voici un exemple simplifié de la façon dont FlashAttention-3 pourrait exploiter le calcul asynchrone :
import torch from torch.cuda.amp import autocast <p>def flash_attention_3(Q, K, V, block_size=256): with autocast(dtype=torch.float8): # En utilisant le FP8 pour le calcul # ... (similaire à la mise en œuvre précédente)</p> <p># Exemple de calcul asynchrone with torch.cuda.stream(torch.cuda.Stream()): # Calculer GEMM de manière asynchrone S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p> <p># Pendant ce temps, sur le flux par défaut: # Préparer pour le calcul softmax</p> <p># Synchroniser les flux torch.cuda.synchronize()</p> <p># Continuer avec le calcul softmax et la sortie # ...</p> return O
Ce code illustre comment FlashAttention-3 pourrait exploiter le calcul asynchrone et la précision FP8. Notez que c’est un exemple simplifié et que la mise en œuvre réelle serait beaucoup plus complexe et spécifique au matériel.
Mettre en œuvre Flash Attention dans vos projets
Si vous êtes enthousiasmé par l’idée d’exploiter Flash Attention dans vos propres projets, vous avez plusieurs options :
- Utiliser des bibliothèques existantes : De nombreuses bibliothèques populaires comme Hugging Face Transformers intègrent désormais des mises en œuvre de Flash Attention. La mise à jour vers la dernière version et l’activation des flags appropriés peut suffire.
- Mise en œuvre personnalisée : Pour plus de contrôle ou pour des cas d’utilisation spécialisés, vous pourriez souhaiter mettre en œuvre Flash Attention vous-même. La bibliothèque xformers fournit une bonne implémentation de référence.
- Optimisations spécifiques au matériel : Si vous travaillez avec un matériel spécifique (par exemple, les GPU NVIDIA H100), vous pourriez vouloir exploiter les fonctionnalités spécifiques au matériel pour une performance maximale.
Voici un exemple de la façon dont vous pourriez utiliser Flash Attention avec la bibliothèque Hugging Face Transformers :
from transformers import AutoModel, AutoConfig <p># Activer Flash Attention config = AutoConfig.from_pretrained("bert-base-uncased") config.use_flash_attention = True</p> <p># Charger le modèle avec Flash Attention model = AutoModel.from_pretrained("bert-base-uncased", config=config)</p> <p># Utiliser le modèle comme d'habitude # ...
Défis et orientations futures
Bien que Flash Attention ait fait des progrès importants dans l’amélioration de l’efficacité des mécanismes d’attention, il existe encore des défis et des domaines de recherche futurs :
- Spécificité du matériel : Les mises en œuvre actuelles sont souvent optimisées pour des architectures GPU spécifiques. Généraliser ces optimisations à travers différents matériels reste un défi.
- Intégration avec d’autres techniques : Combiner Flash Attention avec d’autres techniques d’optimisation comme le élagage, la quantification et la compression de modèle est un domaine de recherche actif.
- Extension à d’autres domaines : Bien que Flash Attention ait montré un grand succès dans le traitement du langage naturel, étendre ses avantages à d’autres domaines comme la vision par ordinateur et les modèles multimodaux est un effort en cours.
- Compréhension théorique : Approfondir notre compréhension théorique de pourquoi Flash Attention fonctionne si bien pourrait conduire à des optimisations encore plus puissantes.
Conclusion
En exploitant de manière astucieuse les hiérarchies de mémoire des GPU et en utilisant des astuces mathématiques, Flash Attention atteint des améliorations substantielles à la fois en vitesse et en utilisation de la mémoire sans sacrifier la précision.
Comme nous l’avons exploré dans cet article, l’impact de Flash Attention s’étend bien au-delà d’une simple technique d’optimisation. Il a permis le développement de modèles plus puissants et plus efficaces.
J'ai passé les cinq dernières années à plonger dans le monde fascinant de l'apprentissage automatique et du deep learning. Ma passion et mon expertise m'ont conduit à contribuer à plus de 50 projets de génie logiciel divers, avec un focus particulier sur l'IA/ML. Ma curiosité continue m'a également attiré vers le traitement automatique des langues, un domaine que je suis impatient d'explorer plus en profondeur.
Découvrir plus


Votre cache KV n’a pas de problème de bits, mais un problème de géométrie.


La course aux armes de l’IA s’intensifie : le partenariat stratégique d’AMD avec OpenAI


Le secret pour un AI plus rapide n’est pas plus de GPUs, mais un réseau plus intelligent


Pourquoi les grands modèles de langage oublient le milieu : dévoiler le point aveugle caché de l’IA


NVIDIA publie une mise à jour corrective pour le problème de surchauffe du pilote GPU


Sapiens : Fondation pour les modèles de vision humaine

