Module 9 — Coût quadratique et attention efficace
L'attention règle le goulot des RNN mais introduit le sien : sa matrice de scores fait . Multipliez la longueur par dix, et la mémoire est multipliée par cent. Cette contrainte pilote toute l'ingénierie des grands modèles depuis 2020 : FlashAttention, fenêtres glissantes, attention éparse, cache quantifié. Ce module explique le calcul du coût, puis passe en revue les remèdes principaux à connaître.
Le calcul du coût, à faire une fois pour toutes
L'attention standard sur une séquence de longueur et une dimension effectue trois grands calculs :
- projections : chacune opérations, coût ;
- scores : produit d'une matrice par , coût ;
- sortie : produit d'une matrice par , coût .
Le terme quadratique en domine dès que , ce qui est le cas usuel. Plus grave, la mémoire occupée par la matrice de scores et par les poids d'attention est par tête et par exemple. Pour jetons, en float16, la matrice pèse environ 128 Mio par tête, à multiplier par le nombre de têtes et par la taille de lot. On atteint vite les limites d'un GPU grand public.
Une règle chiffrée pratique à mémoriser : la mémoire d'un lot de 1 pour l'attention seule à jetons, têtes et float16 est approximativement octets. À , cela fait 512 Mio pour la seule matrice de scores.
def memoire_attention_octets(L: int, h: int, precision_bits: int = 16) -> int:
"""Memoire du tenseur (h, L, L) de scores."""
return (precision_bits // 8) * h * L * L
for L in [1024, 4096, 16384, 65536]:
m = memoire_attention_octets(L, h=16)
print(f"L = {L:6d} memoire scores : {m / 1024**2:8.1f} Mio")
FlashAttention : le même calcul, mieux orchestré
FlashAttention, publié par Dao et al. en 2022, n'invente pas une nouvelle formule : il ré-implémente l'attention exacte de façon à ne jamais matérialiser la matrice entière en mémoire globale du GPU. L'algorithme parcourt la séquence par blocs, calcule des softmax partiels, et fusionne progressivement les résultats en n'utilisant que la mémoire rapide interne (SRAM) de la carte.
Le gain n'est pas en complexité — on fait toujours opérations — mais en débit mémoire, qui est le vrai goulot des GPU modernes. Sur des séquences longues, FlashAttention accélère l'attention d'un facteur 2 à 8 tout en réduisant la mémoire de plusieurs Gio. Depuis PyTorch 2.0, le noyau est activé automatiquement quand on utilise scaled_dot_product_attention avec les bonnes conditions.
import torch
import torch.nn.functional as F
Q = torch.randn(1, 16, 4096, 64, device="cuda", dtype=torch.float16)
K = torch.randn_like(Q)
V = torch.randn_like(Q)
sortie = F.scaled_dot_product_attention(Q, K, V, is_causal=True)
print(sortie.shape) # (1, 16, 4096, 64)
Le drapeau is_causal=True remplace un masque explicite ; il est essentiel de le passer pour que le noyau évite de générer la matrice de masque et applique la causalité directement dans l'algorithme.
Contrairement aux méthodes éparses ci-dessous, FlashAttention calcule exactement la même attention qu'une implémentation naïve. Il n'y a donc pas d'arbitrage qualité contre vitesse : si le matériel le permet, on l'active toujours.
Attention par fenêtre glissante
Une famille de méthodes accepte une approximation pour couper le coût quadratique. La plus simple est l'attention par fenêtre glissante : chaque jeton ne peut regarder que les jetons de part et d'autre, avec typiquement ou . Le coût de calcul et de mémoire redevient linéaire, .
Longformer et Mistral utilisent cette idée. Longformer y ajoute quelques jetons « globaux » qui, eux, voient toute la séquence, pour ne pas perdre d'information de synthèse. Mistral empile plusieurs couches à fenêtre, ce qui construit un contexte effectif plus large : sur couches, un jeton peut indirectement recevoir de l'information à distance , comme dans une convolution profonde.
Le compromis est simple : pour des tâches où les dépendances lointaines sont rares (documents longs mais localement cohérents), la fenêtre suffit. Pour des tâches où un mot doit influencer directement un mot très éloigné (co-références en fiction longue), elle dégrade la qualité.
Attention éparse et modèles hybrides
BigBird, Sparse Transformer et plusieurs variantes proposent des motifs d'éparsité plus riches : combinaison de fenêtres locales, de connexions aléatoires globales, et de jetons de résumé. La théorie montre que ces motifs conservent la capacité d'approximation d'une attention dense sous des hypothèses raisonnables, tout en revenant à un coût linéaire.
Une autre voie, très active depuis 2023, est celle des modèles à espace d'état (Mamba, Mamba-2) qui abandonnent l'attention pour un mécanisme récurrent au coût linéaire, tout en gardant la capacité de gérer les longues dépendances. Ces modèles ne remplacent pas totalement les Transformeurs — ils atteignent des scores comparables à taille égale sur certaines tâches et légèrement inférieurs sur d'autres — mais ils illustrent que la quadratique n'est peut-être pas une fatalité.
Contexte long : cache, quantification et attention groupée
Sur les grands modèles servis en production, le coût dominant n'est plus l'entraînement mais l'inférence à long contexte. Trois techniques se combinent :
- cache clé-valeur (module 7) : évite de recalculer pour les jetons déjà générés ; sa mémoire croît en par exemple ;
- quantification du cache : passage du float16 à l'int8 ou int4, ce qui divise la mémoire du cache par 2 à 4 sans dégrader significativement la sortie ;
- attention à multiples requêtes (MQA) ou groupée (GQA) : plusieurs têtes de se partagent une seule paire , ce qui divise la mémoire du cache par le facteur de partage. LLaMA-2 utilise GQA avec 8 groupes.
À titre indicatif, sur LLaMA-2 70B, un contexte de 4096 jetons occupe un cache d'environ 1,3 Gio par exemple grâce à GQA, contre plus de 10 Gio dans une variante à multiples têtes classique.
Les fournisseurs annoncent régulièrement des contextes de 100 000 ou 1 million de jetons. La courbe de rappel en fonction de la position — le test « aiguille dans une botte de foin » — montre presque toujours que la qualité chute bien avant la longueur maximale. Un contexte publié à 128 000 est souvent réellement fiable jusqu'à 32 000 ou 64 000 selon le modèle. Vérifiez avant d'ancrer un projet dessus.
Sur le fil rouge
Notre tâche de traduction de dates a des séquences de vingt jetons au maximum, donc le coût quadratique est indolore. Nous n'implémenterons ni fenêtre, ni FlashAttention explicite, ni GQA. En revanche, l'appel PyTorch dans le module 10 utilisera scaled_dot_product_attention quand la version l'expose, ce qui active automatiquement les optimisations disponibles.
En résumé
- L'attention standard coûte en calcul et en mémoire pour la matrice de scores ; ce coût quadratique domine dès que les séquences dépassent quelques milliers de jetons.
- FlashAttention garde le calcul exact mais réorganise les accès mémoire pour éviter de matérialiser la matrice , avec un gain de vitesse important sur GPU récents.
- L'attention par fenêtre glissante, l'attention éparse et les modèles à espace d'état obtiennent un coût linéaire au prix d'une approximation, acceptable pour bien des tâches à long contexte.
- En production, le trio cache KV, quantification du cache et attention groupée (GQA) rend possible le service de grands modèles à contexte long avec une mémoire tenable.
Module suivant : rassembler enfin toutes les briques dans un Transformer PyTorch complet, l'entraîner sur la traduction de dates, et visualiser ses cartes d'attention.