Aller au contenu principal

Évolution de l'Attention et Compression du KV Cache

Lors de l'inférence des modèles de langage autorégressifs, la phase de décodage de l'attention est strictement limitée par la bande passante mémoire. En contexte long, le KV Cache croît linéairement avec la séquence et dépasse rapidement la taille des poids statiques du modèle.


1. Panorama architectural : de MHA à GQA et MLA


2. Comparaison des coûts mémoire et de calcul

Pour une dimension cachée dd, nhn_h têtes, dh=d/nhd_h = d / n_h, LL couches et une séquence SS :

2.1 Multi-Head Attention (MHA)

Le MHA standard alloue des têtes Key et Value indépendantes pour chaque tête Query :

extTailleKVCache(MHA)=2imesLimesnhimesdhimesSimesbextelem ext{Taille KV Cache (MHA)} = 2 imes L imes n_h imes d_h imes S imes b_{ ext{elem}}
  • Goulot d'étranglement : Avec nh=128n_h = 128 et S=128extkS = 128 ext{k}, le KV Cache d'une seule requête dépasse 50 Go.

2.2 Grouped-Query Attention (GQA)

Le GQA regroupe nhn_h têtes Query pour partager nkvn_{kv} têtes Key/Value (généralement nkv=8n_{kv} = 8) :

extTailleKVCache(GQA)=2imesLimesnkvimesdhimesSimesbextelem ext{Taille KV Cache (GQA)} = 2 imes L imes n_{kv} imes d_h imes S imes b_{ ext{elem}}
  • Compromis : Compression fixe de nh/nkvn_h / n_{kv} (4imes4 imes8imes8 imes), avec une légère perte d'expressivité dans les contextes denses.

3. Architecture DeepSeek MLA (Multi-Head Latent Attention)

Le MLA introduit une compression conjointe de bas rang, réduisant le KV Cache de plus de 93 % tout en préservant l'expressivité multi-têtes complète.

3.1 Projection de bas rang et vecteur latent

Au lieu de stocker toutes les têtes Key et Value, l'état caché hth_t est projeté dans un espace latent ct(KV)Rdcc_t^{(KV)} \in \mathbb{R}^{d_c} (dcnhimesdhd_c \ll n_h imes d_h) :

ct(KV)=WDKVht,[kt,1(C),,kt,nh(C)]=WUKct(KV),[vt,1(C),,vt,nh(C)]=WUVct(KV)c_t^{(KV)} = W_{DKV} h_t, \quad [k_{t,1}^{(C)}, \dots, k_{t,n_h}^{(C)}] = W_{UK} c_t^{(KV)}, \quad [v_{t,1}^{(C)}, \dots, v_{t,n_h}^{(C)}] = W_{UV} c_t^{(KV)}
  • Absorption matricielle : En VRAM, seul le vecteur latent ct(KV)c_t^{(KV)} est conservé. Les matrices WUVW_{UV} et WUKW_{UK} sont absorbées directement dans les matrices de Query, évitant de matérialiser le Value Cache en mémoire.

3.2 Decoupled RoPE

Les encodages positionnels rotatifs (RoPE) ne pouvant être absorbés dans une projection statique, une tête dédiée kt(R)RdRk_t^{(R)} \in \mathbb{R}^{d_R} transporte les coordonnées positionnelles :

extTailleparToken(MLA)=(dc+dR)imesbextelem ext{Taille par Token (MLA)} = (d_c + d_R) imes b_{ ext{elem}}
ArchitectureCache KV par Token et par Couche (fp16)Empreinte RelativeExpressivité
MHA2imes128imes128imes2=65536extoctets2 imes 128 imes 128 imes 2 = 65\,536 ext{ octets}100%100\% (Référence)Têtes totalement indépendantes
GQA (nkv=8n_{kv}=8)2imes8imes128imes2=4096extoctets2 imes 8 imes 128 imes 2 = 4\,096 ext{ octets}6.25%6.25\%Multi-têtes contraint par groupe
DeepSeek MLA(512+64)imes2=1152extoctets(512 + 64) imes 2 = 1\,152 ext{ octets}1.75%1.75\% (-98.2%)Multi-têtes intégral

4. Attention éparse : SnapKV et PyramidKV

4.1 SnapKV : Fenêtre d'observation

  • SnapKV observe que les têtes d'attention se concentrent sur des empreintes stables. Une petite fenêtre d'observation en fin de prompt suffit pour identifier et évincer 80 % des jetons historiques inactifs avec une dégradation quasi nulle de perplexité (<0.1< 0.1).

4.2 PyramidKV : Allocation hiérarchique

  • Les couches profondes du Transformer nécessitent moins de jetons historiques que les couches superficielles. PyramidKV réduit l'allocation du KV Cache de manière pyramidale, abaissant l'empreinte VRAM de 50 % à 70 %.

5. Recommandations d'ingénierie

  1. Serveurs haute cadence et contextes 100k+ : Privilégier les modèles à architecture MLA native (DeepSeek-V3 / DeepSeek-R1).
  2. Inférence locale sur GPU unique : Sur modèles GQA (Qwen, Llama), utiliser SnapKV ou la quantification 4-bit du KV Cache pour ramener 64k jetons de 16 Go à moins de 4 Go.