Aller au contenu principal

Module 3 — Attention multi-têtes

Une seule tête d'attention produit un unique jeu de poids par jeton, donc une unique manière de mélanger les valeurs. Ce n'est pas assez : une phrase encode plusieurs types de relations en même temps — syntaxiques, sémantiques, positionnelles. Le Transformer résout cela en faisant tourner plusieurs attentions en parallèle, chacune sur un sous-espace différent, avant d'en fusionner les sorties. Ce module explique la mécanique, la lit dans PyTorch, et donne le coût en paramètres exact.

Pourquoi une seule tête ne suffit pas

Dans le module 2, chaque jeton produit une requête, une clé et une valeur en dimension complète dd. Si l'on impose au réseau une seule tête, il doit compresser dans ces vecteurs toutes les relations dont il a besoin. Or, dans « le chien que Marie a vu court », le verbe « court » entretient au moins deux relations différentes avec le reste de la phrase : une relation sujet-verbe avec « chien », et une relation temporelle avec « a vu ». Une seule tête doit choisir laquelle des deux dominera son unique distribution d'attention, ce qu'elle fait mal.

L'attention multi-têtes libère cette contrainte : elle scinde la dimension totale en hh sous-espaces indépendants, laisse chaque tête apprendre ses propres projections, et rassemble ensuite les résultats. Chaque tête peut alors se spécialiser dans un type de relation sans écraser les autres.

La mécanique en une équation, deux étapes

Formellement, avec hh têtes et une dimension par tête dk=d/hd_k = d / h :

tetei=Attention(XWQ(i),XWK(i),XWV(i))\mathrm{tete}_i = \mathrm{Attention}(X W_Q^{(i)},\, X W_K^{(i)},\, X W_V^{(i)}) MultiTete(X)=Concat(tete1,,teteh)WO\mathrm{MultiTete}(X) = \mathrm{Concat}(\mathrm{tete}_1, \ldots, \mathrm{tete}_h)\, W_O

Chaque WQ(i)W_Q^{(i)}, WK(i)W_K^{(i)}, WV(i)W_V^{(i)} projette la dimension complète dd vers dkd_k. Après attention, chaque tête produit une matrice L×dkL \times d_k ; les hh sorties sont concaténées sur la dimension des colonnes pour retrouver la dimension dd, puis passées par une projection de sortie WOW_O apprise. C'est cette dernière projection qui apprend à combiner les différents regards.

En pratique, on n'écrit pas hh multiplications séparées. On empile les projections dans un unique WQW_Q de sortie dd, puis on réorganise la matrice en un tenseur (B,h,L,dk)(B, h, L, d_k). Cela donne exactement hh produits scalaires en parallèle sur le GPU, sans boucle Python.

Ce que différentes têtes apprennent

L'article original visualise les cartes d'attention couche par couche et tête par tête. Un pattern se répète, aujourd'hui bien documenté dans la littérature interprétabilité :

  • certaines têtes s'occupent des relations locales — le mot précédent, le mot suivant — et ressemblent à des filtres de convolution 1D ;
  • d'autres capturent des relations syntaxiques longues, comme un pronom pointant vers son antécédent trente mots plus loin ;
  • d'autres encore agrègent une information positionnelle, ou attirent l'attention vers un jeton spécial comme [CLS] ou [SEP] ;
  • une fraction non négligeable est presque inutilisée — c'est un fait mesuré, sur lequel s'appuient les techniques d'élagage de têtes.

Cette diversité est une conséquence directe de la structure : donner plusieurs têtes à des projections indépendantes crée une pression à la spécialisation.

Le nombre de têtes n'est pas un curseur libre

Les architectures usuelles retiennent 8, 12 ou 16 têtes pour un d=512,768,1024d = 512, 768, 1024. Le rapport d/hd / h (typiquement 64) est plus important que hh lui-même : c'est la dimension de chaque tête, et elle doit rester assez large pour que le produit scalaire ait un sens statistique.

Le coût en paramètres se calcule à la main

C'est un exercice classique et il tombe souvent à l'examen. Pour une couche multi-têtes avec dimension dd et hh têtes de dimension dk=d/hd_k = d/h chacune, on a :

  • trois projections d'entrée WQ,WK,WVW_Q, W_K, W_V de taille d×dd \times d chacune, soit 3d23 d^2 paramètres ;
  • une projection de sortie WOW_O de taille d×dd \times d, soit d2d^2 paramètres ;
  • au total 4d24 d^2 paramètres par couche multi-têtes.

Autrement dit, le nombre de têtes ne change pas le nombre de paramètres, tant que dd reste fixe. Passer de 8 à 16 têtes divise seulement dkd_k par deux, ce qui augmente le nombre de sous-espaces sans coûter un paramètre de plus.

Pour un modèle avec d=512d = 512, une couche multi-têtes utilise donc 4×5122=10485764 \times 512^2 = 1\,048\,576 paramètres. C'est un point de repère utile pour dimensionner un modèle. Le bloc feed-forward qui suit chaque couche d'attention en utilise en général deux à quatre fois plus (module 5), et c'est lui qui domine le budget total.

Implémentation en PyTorch, sans boucle

Voici la version qu'on gardera pour le fil rouge :

import torch
import torch.nn as nn
import torch.nn.functional as F

class AttentionMultiTetes(nn.Module):
def __init__(self, d_modele: int, n_tetes: int, dropout: float = 0.1):
super().__init__()
assert d_modele % n_tetes == 0, "d_modele doit etre divisible par n_tetes"
self.h = n_tetes
self.d_k = d_modele // n_tetes
self.w_q = nn.Linear(d_modele, d_modele, bias=False)
self.w_k = nn.Linear(d_modele, d_modele, bias=False)
self.w_v = nn.Linear(d_modele, d_modele, bias=False)
self.w_o = nn.Linear(d_modele, d_modele, bias=False)
self.dropout = nn.Dropout(dropout)

def forward(self, q_in, k_in, v_in, masque=None):
B, L_q, D = q_in.shape
L_k = k_in.shape[1]

# (B, L, D) -> (B, h, L, d_k)
def decouper(x, L):
return x.view(B, L, self.h, self.d_k).transpose(1, 2)

Q = decouper(self.w_q(q_in), L_q)
K = decouper(self.w_k(k_in), L_k)
V = decouper(self.w_v(v_in), L_k)

scores = Q @ K.transpose(-2, -1) / (self.d_k ** 0.5)
if masque is not None:
scores = scores.masked_fill(masque == 0, float("-inf"))
poids = self.dropout(F.softmax(scores, dim=-1))

sortie = poids @ V # (B, h, L_q, d_k)
sortie = sortie.transpose(1, 2).contiguous().view(B, L_q, D)
return self.w_o(sortie), poids

couche = AttentionMultiTetes(d_modele=64, n_tetes=8)
x = torch.randn(2, 12, 64)
sortie, poids = couche(x, x, x)
print(sortie.shape, poids.shape) # (2, 12, 64) (2, 8, 12, 12)

Trois points méritent l'attention. Le assert garantit que dd est bien divisible par hh, sans quoi la vue échouerait avec un message peu lisible. Le transpose(1, 2) puis contiguous() réorganise la mémoire avant le view final ; sans contiguous(), PyTorch lève une exception. Le dropout sur les poids d'attention est le choix original de l'article ; il aide à la régularisation sans changer la forme du calcul.

Séparer qq, kk, vv dans la signature

Notre forward prend trois entrées, pas une seule. C'est indispensable dès qu'on veut faire de l'attention croisée : la requête vient du décodeur, les clés et valeurs viennent de l'encodeur. Pour l'auto-attention, on appelle simplement couche(x, x, x).

En résumé

  • L'attention multi-têtes découpe la dimension dd en hh sous-espaces indépendants de dimension dk=d/hd_k = d/h, ce qui permet à chaque tête de se spécialiser.
  • Les têtes se calculent en parallèle par un simple view puis transpose, sans boucle Python : c'est ce qui rend la couche efficace sur GPU.
  • Le coût est 4d24 d^2 paramètres — trois projections d'entrée et une projection de sortie — indépendant du nombre de têtes.
  • Les cartes d'attention observées confirment la spécialisation : têtes locales, syntaxiques, positionnelles, et une part significative de têtes peu utilisées.

Module suivant : puisque l'attention ignore l'ordre des jetons, il faut le lui réinjecter par un encodage de position, absolu ou rotatif.