Aller au contenu principal

Module 2 — Requête, clé, valeur : l'attention pas à pas

L'attention est présentée dans la plupart des tutoriels par une formule compacte, ce qui la rend mystérieuse. Ce module fait l'inverse : on prend trois jetons du fil rouge, on calcule tout à la main, et on écrit la couche en une vingtaine de lignes de PyTorch. À la fin du module, la formule ne cache plus rien.

Trois vecteurs par jeton, trois rôles distincts

Le point de départ de l'attention est une intuition d'annuaire. Chaque jeton de la séquence joue trois rôles simultanés, portés par trois vecteurs qu'on calcule à partir de son plongement de départ :

  • la requête qiq_i décrit ce que ce jeton cherche à savoir ;
  • la clé kik_i décrit ce que ce jeton peut fournir aux autres ;
  • la valeur viv_i est le contenu effectif qu'il transmet quand on l'interroge.

Pour comparer deux jetons, on confronte la requête du premier aux clés de tous les autres. Plus la requête et une clé pointent dans une direction proche, plus le score est élevé, et plus la valeur associée pèsera dans la moyenne finale.

Formellement, si XRL×dX \in \mathbb{R}^{L \times d} est la matrice des plongements de LL jetons en dimension dd, on calcule :

Q=XWQ,K=XWK,V=XWVQ = X W_Q, \quad K = X W_K, \quad V = X W_V

WQW_Q, WKW_K et WVW_V sont trois matrices de projection apprises. Les trois vecteurs par jeton ne sortent pas de nulle part : ils sont trois vues linéaires du même plongement d'entrée, ajustées par la rétropropagation.

La formule complète, décortiquée

La sortie de l'attention pour toute la séquence s'écrit :

Attention(Q,K,V)=softmax ⁣(QKdk)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\!\left( \frac{Q K^\top}{\sqrt{d_k}} \right) V

Il faut la lire dans l'ordre.

D'abord, QKQ K^\top est une matrice L×LL \times L où l'entrée (i,j)(i, j) est le produit scalaire de qiq_i et kjk_j. C'est un score brut de similarité entre la requête du jeton ii et la clé du jeton jj.

Ensuite, la division par dk\sqrt{d_k} — où dkd_k est la dimension des clés — évite que les scores explosent quand la dimension grandit. Un produit scalaire de deux vecteurs aléatoires de dimension dkd_k a une variance proportionnelle à dkd_k ; sans mise à l'échelle, le softmax qui suit devient presque déterministe et sature ses gradients. La racine carrée compense exactement.

Puis le softmax appliqué ligne par ligne transforme ces scores en distributions de probabilité : la ligne ii donne pour chaque jj le poids d'attention que ii accorde à jj. La somme de chaque ligne fait 1.

Enfin, la multiplication par VV fait la moyenne pondérée des valeurs par ces poids. La ligne ii du résultat est un mélange, dosé par le softmax, des valeurs de tous les jetons de la séquence.

Softmax appliqué ligne par ligne, jamais colonne par colonne

C'est une source classique d'erreur d'implémentation. Le softmax doit être appliqué sur la dernière dimension de la matrice de scores L×LL \times L, pour que chaque ligne — c'est-à-dire chaque requête — devienne une distribution qui somme à 1. En PyTorch, cela s'écrit torch.softmax(scores, dim=-1).

Un calcul numérique complet sur trois jetons

Prenons trois jetons du fil rouge, très simplifiés, en dimension d=dk=2d = d_k = 2 :

import numpy as np

# Trois jetons : "3", "mars", "2026", plongements de dimension 2.
X = np.array([
[1.0, 0.0], # "3"
[0.0, 1.0], # "mars"
[1.0, 1.0], # "2026"
])

# Projections apprises, fixees ici pour l'exemple.
W_Q = np.array([[1.0, 0.0], [0.0, 1.0]])
W_K = np.array([[1.0, 0.0], [0.0, 1.0]])
W_V = np.array([[0.0, 1.0], [1.0, 0.0]])

Q = X @ W_Q
K = X @ W_K
V = X @ W_V
d_k = 2

scores = Q @ K.T / np.sqrt(d_k)
poids = np.exp(scores - scores.max(axis=1, keepdims=True))
poids /= poids.sum(axis=1, keepdims=True)
sortie = poids @ V

print("scores :\n", scores.round(2))
print("poids :\n", poids.round(2))
print("sortie :\n", sortie.round(2))

Avec ces projections identiques pour QQ et KK, la matrice de scores est symétrique. On y lit que « 3 » et « 2026 » s'aligne mieux entre eux (score 1/21/\sqrt{2}) qu'avec « mars ». Après softmax et multiplication par VV, chaque jeton reçoit un mélange des trois valeurs pondéré selon ces affinités. C'est exactement ce que fait la couche MultiheadAttention de PyTorch, avec en plus un traitement multi-têtes que le module 3 introduira.

Auto-attention : la source et la cible sont la même séquence

Dans l'exemple ci-dessus, les trois vecteurs QQ, KK et VV proviennent tous de la même matrice XX. On parle alors d'auto-attention : la séquence s'interroge elle-même. Chaque jeton peut consulter tous les autres, y compris lui-même, et se recontextualiser en fonction de son voisinage — parfois immédiat, parfois lointain.

Cette liberté fait la force de la couche. Un modèle récurrent doit relayer une information de mot en mot ; l'auto-attention y accède en un seul saut, avec un coût de calcul indépendant de la distance entre les jetons. C'est pourquoi les Transformeurs modélisent si bien les dépendances longues.

Quand QQ vient d'une séquence et que K,VK, V viennent d'une autre — typiquement, le décodeur consulte l'encodeur — on parle d'attention croisée. C'est le mécanisme qui rend possible la traduction, et il sera étudié au module 8.

Une implémentation courte en PyTorch

Voici la couche d'attention à une tête, exactement comme on l'écrira dans notre Transformer maison :

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

class AttentionSimple(nn.Module):
def __init__(self, d_modele: int):
super().__init__()
self.d_k = d_modele
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)

def forward(self, x: torch.Tensor, masque: torch.Tensor | None = None):
# x : (B, L, D)
Q, K, V = self.w_q(x), self.w_k(x), self.w_v(x)
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
if masque is not None:
scores = scores.masked_fill(masque == 0, float("-inf"))
poids = F.softmax(scores, dim=-1)
return torch.matmul(poids, V), poids

couche = AttentionSimple(d_modele=32)
x = torch.randn(4, 10, 32)
sortie, poids = couche(x)
print(sortie.shape, poids.shape) # (4, 10, 32) (4, 10, 10)

Le paramètre masque sera utilisé au module 7 pour empêcher un décodeur autorégressif de regarder vers le futur. Ici, on l'introduit dès maintenant, car sa gestion propre évite bien des bogues plus tard.

L'ordre -inf puis softmax est capital

Un mauvais réflexe consiste à mettre les scores masqués à zéro avant le softmax. Mais e0=1e^0 = 1 : le jeton masqué reçoit alors un poids non nul. Il faut mettre les positions interdites à -inf avant le softmax, ce qui les envoie à un poids exactement nul.

En résumé

  • Chaque jeton porte trois vecteurs, requête, clé et valeur, obtenus par trois projections linéaires apprises à partir du plongement d'entrée.
  • Les scores QKQ K^\top sont divisés par dk\sqrt{d_k} pour éviter la saturation du softmax quand la dimension des clés grandit.
  • Le softmax est ligne par ligne ; la sortie est la moyenne des valeurs pondérée par ces poids, calculée en une seule multiplication matricielle.
  • L'auto-attention partage Q,K,VQ, K, V sur la même séquence, tandis que l'attention croisée relie deux séquences (décodeur regardant l'encodeur).

Module suivant : pourquoi une seule tête d'attention ne suffit pas, et comment plusieurs têtes en parallèle capturent des relations de nature différente.