Cours LLM Inférence

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.

Écouter ce cours (14 min)

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

Comprendre
Pourquoi un modèle de langage génère du texte token par token et pourquoi le contexte doit être conservé.
Maîtriser
Les formules de l’attention causale, du calcul de queries, keys, values et du softmax d’attention.
Expliquer
Pourquoi le KV cache est lu à chaque itération pendant la phase de décodage.
Analyser
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.
Note : ici, nous nous concentrons surtout sur l’inférence, c’est-à-dire la phase où le modèle génère du texte après avoir été entraîné.

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.

Idée clé : le KV cache est la mémoire des tokens précédents. Il évite de recalculer entièrement tout le contexte à chaque nouveau token.

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 :

$$ p(x_{1:T}) = \prod_{t=1}^{T} p(x_t \mid x_{1:t-1}) $$

Cela signifie que chaque token est prédit à partir des tokens précédents :

$$ x_t \sim p(\cdot \mid x_{1:t-1}) $$

Pendant la génération, la séquence grandit donc progressivement :

$$ x_1, x_2, x_3, \dots, x_t $$
Conséquence : à chaque itération, le modèle doit prendre en compte un contexte plus long que lors de l’itération précédente.

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 :

$$ Q = X W_Q $$ $$ K = X W_K $$ $$ V = X W_V $$

L’attention standard est définie par :

$$ \operatorname{Attention}(Q, K, V) = \operatorname{softmax} \left( \frac{QK^\top}{\sqrt{d_k}} \right) V $$

Dans un modèle de langage causal, le token courant ne doit pas voir les tokens futurs. On applique donc un masque causal :

$$ A_{ij} = \begin{cases} \frac{q_i^\top k_j}{\sqrt{d_k}}, & \text{si } j \le i \\ -\infty, & \text{si } j > i \end{cases} $$

Puis :

$$ \alpha_{ij} = \frac{\exp(A_{ij})}{\sum_{m=1}^{i} \exp(A_{im})} $$

Pour un token courant $t$, les poids d’attention sur les tokens précédents sont :

$$ \alpha_{tj} = \frac{ \exp\left( \frac{q_t^\top k_j}{\sqrt{d_k}} \right) } { \sum_{i=1}^{t} \exp\left( \frac{q_t^\top k_i}{\sqrt{d_k}} \right) } $$

La sortie d’attention pour le token $t$ est ensuite :

$$ o_t = \sum_{j=1}^{t} \alpha_{tj} v_j $$
Point important : la sortie $o_t$ dépend de toutes les valeurs précédentes $v_1, \dots, v_t$, mais aussi de la query courante $q_t$.

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 :

$$ q_t = h_t W_Q $$ $$ k_t = h_t W_K $$ $$ v_t = h_t W_V $$

Ensuite, il met à jour le cache :

$$ K_{\text{cache}}^{(t)} = \left[ K_{\text{cache}}^{(t-1)} ; k_t \right] $$ $$ V_{\text{cache}}^{(t)} = \left[ V_{\text{cache}}^{(t-1)} ; v_t \right] $$

Puis il calcule l’attention du token courant en utilisant la query courante et tout le cache :

$$ o_t = \operatorname{Attention} \left( q_t, K_{\text{cache}}^{(t)}, V_{\text{cache}}^{(t)} \right) $$

Pour une couche $\ell$, on peut écrire :

$$ K_{\text{cache}, t}^{(\ell)} = \operatorname{Concat} \left( K_{\text{cache}, t-1}^{(\ell)}, h_t^{(\ell-1)} W_K^{(\ell)} \right) $$ $$ V_{\text{cache}, t}^{(\ell)} = \operatorname{Concat} \left( V_{\text{cache}, t-1}^{(\ell)}, h_t^{(\ell-1)} W_V^{(\ell)} \right) $$
Rappel : chaque couche possède généralement son propre KV cache. Il ne s’agit donc pas d’un seul cache global, mais d’un ensemble de caches par couche.

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 :

$$ q_t $$

Au step suivant, elle devient :

$$ q_{t+1} $$

Or les poids d’attention dépendent directement de cette query :

$$ \alpha_{tj} = \operatorname{softmax}_j \left( \frac{q_t^\top k_j}{\sqrt{d_k}} \right) $$

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.

$$ o_t = \sum_{j=1}^{t} \alpha_{tj} v_j $$

À l’itération suivante :

$$ o_{t+1} = \sum_{j=1}^{t+1} \alpha_{t+1,j} v_j $$

Les poids $\alpha_{tj}$ et $\alpha_{t+1,j}$ ne sont pas identiques, car ils dépendent de queries différentes.

Conclusion : on ne peut pas simplement réutiliser la sortie d’attention du token précédent. Il faut recalculer l’attention du nouveau token sur tout le passé.

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 :

$$ M_{KV} = 2 \times B \times L \times T \times H_{KV} \times d_h \times b $$

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
$$ M_{KV} = 2 \times 1 \times 32 \times 32768 \times 8 \times 128 \times 2 $$
$$ M_{KV} \approx 4.29 \times 10^9 \text{ octets} $$

Cela représente environ 4 Go de mémoire pour une seule requête.

Problème : pendant la génération token par token, le modèle passe souvent plus de temps à lire le KV cache qu’à effectuer des opérations arithmétiques. L’inférence devient alors limitée par la bande passante mémoire.

Complexité temporelle

Pour un token de décodage, le coût de lecture/calcul de l’attention est proportionnel à :

$$ O(T \cdot d) $$

Pour $T$ tokens générés, le coût cumulé devient :

$$ O(T^2 \cdot d) $$

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 :

$$ H_{KV} = H_Q $$

Avec GQA, si les query heads sont regroupées :

$$ H_{KV} < H_Q $$

Avec MQA :

$$ H_{KV} = 1 $$

La mémoire du cache est donc réduite proportionnellement.

Formule avec sliding window

Si l’on ne conserve que les $W$ derniers tokens :

$$ M_{KV}^{\text{window}} = 2 \times B \times L \times W \times H_{KV} \times d_h \times b $$

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é :

$$ h_t = f(h_{t-1}, x_t) $$

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 :

$$ s_t = A s_{t-1} + B x_t $$ $$ y_t = C s_t $$

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é :

$$ S_t = S_{t-1} + \phi(k_t) v_t^\top $$

Puis la sortie peut être calculée approximativement par :

$$ o_t = \frac{ \phi(q_t)^\top S_t } { \phi(q_t)^\top z_t } $$

où $z_t$ est un terme de normalisation.

Compromis : ces architectures peuvent réduire le coût de génération, mais elles peuvent perdre une partie de la capacité d’accès précis au contexte offerte par l’attention complète du Transformer.

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

Génération
Le modèle prédit un token à la fois selon $p(x_t \mid x_{1:t-1})$.
Attention
Le token courant utilise sa query $q_t$ pour interroger les anciennes clés $k_j$.
KV cache
Il conserve les anciens $K$ et $V$ pour éviter de recalculer tout le passé.
Relecture
Le cache est relu car la query change à chaque token, donc les poids d’attention changent.
Coût
Le cache peut devenir volumineux : $M_{KV} = 2BLTH_{KV}d_hb$.
Optimisations
GQA, MQA, quantification, PagedAttention, sliding window, sparse attention, speculative decoding.
Phrase à retenir : le KV cache est la mémoire contextuelle du Transformer pendant la génération. Le modèle ne relit pas le cache par inefficacité, mais parce que chaque nouveau token doit recalculer sa propre attention sur tout le passé.