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 . 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 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 têtes et une dimension par tête :
Chaque , , projette la dimension complète vers . Après attention, chaque tête produit une matrice ; les sorties sont concaténées sur la dimension des colonnes pour retrouver la dimension , puis passées par une projection de sortie apprise. C'est cette dernière projection qui apprend à combiner les différents regards.
En pratique, on n'écrit pas multiplications séparées. On empile les projections dans un unique de sortie , puis on réorganise la matrice en un tenseur . Cela donne exactement 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.
Les architectures usuelles retiennent 8, 12 ou 16 têtes pour un . Le rapport (typiquement 64) est plus important que 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 et têtes de dimension chacune, on a :
- trois projections d'entrée de taille chacune, soit paramètres ;
- une projection de sortie de taille , soit paramètres ;
- au total paramètres par couche multi-têtes.
Autrement dit, le nombre de têtes ne change pas le nombre de paramètres, tant que reste fixe. Passer de 8 à 16 têtes divise seulement par deux, ce qui augmente le nombre de sous-espaces sans coûter un paramètre de plus.
Pour un modèle avec , une couche multi-têtes utilise donc 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 est bien divisible par , 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.
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 en sous-espaces indépendants de dimension , ce qui permet à chaque tête de se spécialiser.
- Les têtes se calculent en parallèle par un simple
viewpuistranspose, sans boucle Python : c'est ce qui rend la couche efficace sur GPU. - Le coût est 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.