Pourquoi un LLM relit son KV cache à chaque itération ?
Ce cours explique le rôle du KV cache dans les Transformers autoregressifs, pourquoi le modèle doit relire les anciennes clés et valeurs à chaque génération de token, quelles sont les formules mathématiques associées, et quelles optimisations permettent de réduire le coût mémoire et bande passante.
Version audio générée localement (Pocket TTS, Kyutai, voix Estelle). Les formules et tableaux sont à retrouver sur cette page.
1. Objectifs pédagogiques
Pourquoi un modèle de langage génère du texte token par token et pourquoi le contexte doit être conservé.
Les formules de l’attention causale, du calcul de queries, keys, values et du softmax d’attention.
Pourquoi le KV cache est lu à chaque itération pendant la phase de décodage.
Le coût mémoire du KV cache et les principales optimisations utilisées en production.
2. Prérequis
Ce cours suppose une connaissance de base des réseaux de neurones et des Transformers. Les notions suivantes sont utiles :
- produit matriciel ;
- softmax ;
- embeddings de tokens ;
- inférence d’un modèle de langage ;
- différence entre entraînement et génération.
3. Intuition générale
Un LLM autoregressif produit un texte mot par mot, ou plus exactement token par token. À chaque étape, il doit choisir le token suivant en fonction de tout ce qui a été généré avant.
Exemple :
Entrée : "Le chat dort sur le"Objectif : prédire "canapé"
Pour prédire correctement le prochain token, le modèle doit pouvoir regarder les tokens précédents. Dans un Transformer, cette capacité est réalisée par le mécanisme d’attention.
4. Génération autoregressive
Un modèle de langage autoregressif factorise la probabilité d’une séquence comme un produit de probabilités conditionnelles :
Cela signifie que chaque token est prédit à partir des tokens précédents :
Pendant la génération, la séquence grandit donc progressivement :
5. Attention causale
Dans un Transformer, chaque token est transformé en trois vecteurs :
- Query : ce que le token cherche ;
- Key : ce que le token contient comme information ;
- Value : ce que le token peut transmettre.
Pour une couche donnée, on calcule :
L’attention standard est définie par :
Dans un modèle de langage causal, le token courant ne doit pas voir les tokens futurs. On applique donc un masque causal :
Puis :
Pour un token courant $t$, les poids d’attention sur les tokens précédents sont :
La sortie d’attention pour le token $t$ est ensuite :
6. Fonctionnement du KV cache
Pendant la génération, les tokens précédents ne changent pas. Leurs représentations de clés et de valeurs peuvent donc être calculées une seule fois, puis conservées dans un cache.
À l’étape $t$, le modèle calcule seulement les vecteurs du nouveau token :
Ensuite, il met à jour le cache :
Puis il calcule l’attention du token courant en utilisant la query courante et tout le cache :
Pour une couche $\ell$, on peut écrire :
7. Pourquoi relire le KV cache à chaque token ?
La raison principale est que la query change à chaque token généré.
Au step $t$, la query est :
Au step suivant, elle devient :
Or les poids d’attention dépendent directement de cette query :
Donc même si les anciennes clés $k_j$ et valeurs $v_j$ ne changent pas, la manière dont le nouveau token doit les pondérer change.
À l’itération suivante :
Les poids $\alpha_{tj}$ et $\alpha_{t+1,j}$ ne sont pas identiques, car ils dépendent de queries différentes.
Analogie
Imagine que tu écris une phrase mot par mot. À chaque nouveau mot, tu dois relire les mots précédents pour savoir si la suite est cohérente. Le KV cache est le carnet dans lequel le modèle garde les mots précédents. Il ne réécrit pas tout, mais il doit relire le carnet à chaque nouveau mot.
8. Coût mémoire et bande passante
Le KV cache évite de recalculer les anciens tokens, mais il doit être stocké puis lu à chaque étape de génération.
La taille mémoire du KV cache peut être estimée par la formule suivante :
Où :
| Symbole | Signification |
|---|---|
| $2$ | car on stocke $K$ et $V$ |
| $B$ | taille du batch |
| $L$ | nombre de couches |
| $T$ | longueur de séquence |
| $H_{KV}$ | nombre de têtes de clés/valeurs |
| $d_h$ | dimension par tête |
| $b$ | nombre d’octets par valeur |
Exemple numérique
Supposons :
- $B = 1$
- $L = 32$ couches
- $T = 32768$ tokens
- $H_{KV} = 8$ têtes KV
- $d_h = 128$
- $b = 2$ octets, par exemple FP16/BF16
Cela représente environ 4 Go de mémoire pour une seule requête.
Complexité temporelle
Pour un token de décodage, le coût de lecture/calcul de l’attention est proportionnel à :
Pour $T$ tokens générés, le coût cumulé devient :
C’est pourquoi les longues séquences sont coûteuses en inférence.
9. Optimisations modernes
Plusieurs techniques permettent de réduire le coût du KV cache.
| Optimisation | Idée principale | Effet |
|---|---|---|
| GQA Grouped Query Attention |
Plusieurs query heads partagent un même groupe de têtes KV. | Réduit $H_{KV}$ et donc la taille du cache. |
| MQA Multi-Query Attention |
Une seule tête KV est partagée par toutes les query heads. | Réduction mémoire très forte, parfois au détriment de la qualité. |
| Quantification du KV cache | Stocker les clés/valeurs en INT8, INT4, etc. | Réduit $b$, donc la mémoire et la bande passante. |
| PagedAttention | Découper le KV cache en pages mémoire, comme la mémoire virtuelle. | Réduit la fragmentation et améliore le batching. |
| Sliding window attention | Ne regarder que les $W$ derniers tokens. | Remplace $O(T)$ par $O(W)$ pour l’attention locale. |
| Sparse attention | Ne calculer l’attention que sur certains tokens sélectionnés. | Réduit le nombre de clés/valeurs réellement lues. |
| Token eviction | Supprimer ou compresser les anciens tokens peu importants. | Réduit la taille effective du cache. |
| Speculative decoding | Utiliser un petit modèle pour proposer plusieurs tokens, vérifiés par le grand modèle. | Augmente le débit de génération. |
Formule avec GQA / MQA
Dans une attention multi-têtes classique :
Avec GQA, si les query heads sont regroupées :
Avec MQA :
La mémoire du cache est donc réduite proportionnellement.
Formule avec sliding window
Si l’on ne conserve que les $W$ derniers tokens :
Pour $W \ll T$, le gain mémoire est important.
10. Alternatives architecturales
Le Transformer classique n’est pas la seule architecture possible. Certaines architectures cherchent à éviter la relecture complète du passé.
10.1 Réseaux récurrents
Dans un RNN, le passé est compressé dans un état caché :
Avantages :
- pas besoin de stocker tous les anciens tokens ;
- mémoire constante en génération.
Inconvénients :
- l’information ancienne peut être perdue ;
- accès moins précis à un token spécifique du passé.
10.2 State Space Models
Les modèles de type State Space Model, comme certaines architectures modernes à état récurrent, maintiennent un état latent mis à jour à chaque étape.
Schématiquement :
L’objectif est d’obtenir une génération efficace tout en conservant une mémoire longue.
10.3 Attention linéaire
Certaines approches remplacent le softmax par une forme linéaire ou approximée. On peut alors maintenir un état agrégé :
Puis la sortie peut être calculée approximativement par :
où $z_t$ est un terme de normalisation.
11. Quiz
Question 1 : Que stocke le KV cache ?
Le KV cache stocke les anciennes clés $K$ et valeurs $V$ calculées pour les tokens précédents. Il ne stocke généralement pas les queries des anciens tokens, car seule la query du token courant est nécessaire pendant le décodage.
Question 2 : Pourquoi ne recalcule-t-on pas tous les anciens tokens à chaque étape ?
Parce que les anciens tokens ne changent pas. Leurs clés et valeurs peuvent donc être calculées une fois puis réutilisées. Cela évite un recalcul coûteux de toute la séquence à chaque nouveau token.
Question 3 : Pourquoi faut-il relire le cache même si les anciennes valeurs ne changent pas ?
Parce que la query du nouveau token change à chaque itération. Les poids d’attention dépendent de cette query :
$$ \alpha_{tj} = \operatorname{softmax}_j \left( \frac{q_t^\top k_j}{\sqrt{d_k}} \right) $$Donc la combinaison des valeurs précédentes doit être recalculée pour chaque nouveau token.
Question 4 : Quelle est la formule de taille mémoire du KV cache ?
Pour un batch $B$, $L$ couches, une séquence de longueur $T$, $H_{KV}$ têtes KV, une dimension par tête $d_h$, et $b$ octets par valeur :
$$ M_{KV} = 2 \times B \times L \times T \times H_{KV} \times d_h \times b $$Question 5 : Pourquoi la génération est-elle souvent limitée par la bande passante ?
Parce qu’à chaque token généré, il faut lire une grande quantité de données depuis le KV cache. Le calcul arithmétique est souvent faible comparé au volume de données à déplacer depuis la mémoire.
Question 6 : Le KV cache est-il obligatoire pour tous les modèles de langage ?
Non. Il est essentiel pour les Transformers autoregressifs classiques avec attention complète. D’autres architectures, comme les RNN, certains State Space Models ou des formes d’attention linéaire, utilisent un état compressé plutôt qu’un cache complet de toutes les anciennes clés/valeurs.
12. Résumé final
Le modèle prédit un token à la fois selon $p(x_t \mid x_{1:t-1})$.
Le token courant utilise sa query $q_t$ pour interroger les anciennes clés $k_j$.
Il conserve les anciens $K$ et $V$ pour éviter de recalculer tout le passé.
Le cache est relu car la query change à chaque token, donc les poids d’attention changent.
Le cache peut devenir volumineux : $M_{KV} = 2BLTH_{KV}d_hb$.
GQA, MQA, quantification, PagedAttention, sliding window, sparse attention, speculative decoding.