Aller au contenu principal

Module 10 — Implémentation d'un bloc Transformer complet

Il est temps d'assembler. Ce module rassemble l'attention multi-têtes du module 3, l'encodage sinusoïdal du module 4, les résidus et la norme du module 5, l'encodeur du module 6, le décodeur du module 7 et l'attention croisée du module 8 dans un unique Transformer encodeur-décodeur d'environ 200 lignes de PyTorch. On l'entraîne sur la traduction de dates française vers ISO, et l'on visualise ce qu'il apprend à regarder.

Le jeu de données jouet

La cible pédagogique impose deux qualités : simple à générer et lisible dans les cartes d'attention. La traduction de dates coche les deux.

import random
from datetime import date, timedelta

MOIS_FR = ["janvier", "fevrier", "mars", "avril", "mai", "juin",
"juillet", "aout", "septembre", "octobre", "novembre", "decembre"]

def echantillon_date():
d = date(2020, 1, 1) + timedelta(days=random.randint(0, 365 * 6))
source = f"{d.day} {MOIS_FR[d.month - 1]} {d.year}"
cible = d.isoformat()
return source, cible

random.seed(0)
for _ in range(3):
print(echantillon_date())
# ('27 juin 2020', '2020-06-27')
# ('2 aout 2023', '2023-08-02')
# ('19 mars 2025', '2025-03-19')

On garde 10 000 exemples pour l'entraînement, 1 000 pour la validation. Le vocabulaire source contient les chiffres, les mois écrits en toutes lettres et l'espace ; le vocabulaire cible contient les chiffres et le tiret. Deux jetons spéciaux, [BOS] et [EOS], encadrent la cible.

Assemblage du modèle

Les briques des modules précédents s'enchaînent proprement. On réutilise AttentionMultiTetes du module 3 et encodage_sinusoidal du module 4 sans les redéfinir.

import torch
import torch.nn as nn

class Transformer(nn.Module):
def __init__(self, taille_src, taille_cbl, d=64, h=8, d_ff=256,
n_encodeurs=2, n_decodeurs=2, dropout=0.1, L_max=32):
super().__init__()
self.d = d
self.plongement_src = nn.Embedding(taille_src, d)
self.plongement_cbl = nn.Embedding(taille_cbl, d)
self.register_buffer("pe", encodage_sinusoidal(L_max, d))

self.encodeurs = nn.ModuleList(
[CoucheEncodeur(d, h, d_ff, dropout) for _ in range(n_encodeurs)]
)
self.decodeurs = nn.ModuleList(
[CoucheDecodeur(d, h, d_ff, dropout) for _ in range(n_decodeurs)]
)
self.norme_finale = nn.LayerNorm(d)
self.sortie = nn.Linear(d, taille_cbl, bias=False)

def encoder(self, src, masque_src=None):
L = src.size(1)
x = self.plongement_src(src) + self.pe[:L].unsqueeze(0)
for couche in self.encodeurs:
x = couche(x, masque=masque_src)
return x

def decoder(self, cbl, memoire, masque_causal=None, masque_src=None):
L = cbl.size(1)
x = self.plongement_cbl(cbl) + self.pe[:L].unsqueeze(0)
for couche in self.decodeurs:
x = couche(x, memoire, masque_causal, masque_src)
return self.sortie(self.norme_finale(x))

def forward(self, src, cbl, masque_causal=None, masque_src=None):
memoire = self.encoder(src, masque_src)
return self.decoder(cbl, memoire, masque_causal, masque_src)

Le nombre total de paramètres, pour d=64,dff=256d = 64, d_{\mathrm{ff}} = 256 et deux couches par pile, tourne autour de 200 000. C'est délibérément petit : le but n'est pas la performance mais la visualisation. Un modèle plus grand cache ses régularités dans un grand nombre de têtes ; un petit modèle est forcé de faire tenir la structure de la date dans un ou deux motifs d'attention lisibles.

Boucle d'entraînement, à taille de portable

L'entraînement tient sur CPU en quelques minutes.

import torch.optim as optim

modele = Transformer(taille_src=len(vocab_src), taille_cbl=len(vocab_cbl))
optimiseur = optim.AdamW(modele.parameters(), lr=3e-4, weight_decay=0.01)
perte = nn.CrossEntropyLoss(ignore_index=vocab_cbl["[PAD]"])

for epoque in range(30):
modele.train()
perte_totale = 0.0
for src, cbl_entree, cbl_sortie in charger_lots(taille_lot=64):
L_cbl = cbl_entree.size(1)
masque_causal = masque_causal_grille(L_cbl) # module 7, adapte

logits = modele(src, cbl_entree, masque_causal=masque_causal)
p = perte(logits.reshape(-1, logits.size(-1)), cbl_sortie.reshape(-1))
optimiseur.zero_grad()
p.backward()
torch.nn.utils.clip_grad_norm_(modele.parameters(), 1.0)
optimiseur.step()
perte_totale += p.item()

print(f"epoque {epoque:2d} perte moyenne : {perte_totale:.3f}")

Trois détails valent d'être notés. ignore_index=[PAD] évite que les jetons de bourrage — nécessaires pour aligner les longueurs dans un lot — contribuent à la perte. Le clipping du gradient à 1.0 empêche une explosion rare mais possible sur les premières étapes. Enfin, cbl_entree et cbl_sortie sont décalées d'un jeton : à l'entrée du décodeur on met [BOS] 2 0 2 6 - 0 3 - 0 3, en sortie attendue 2 0 2 6 - 0 3 - 0 3 [EOS]. C'est la prédiction du jeton suivant du module 7.

Ne pas oublier le décalage cible

Un décodeur qui reçoit la même séquence en entrée et en sortie attendue n'apprend rien : il lui suffit de recopier. Le décalage d'un jeton, avec [BOS] en tête et [EOS] en fin, est la seule façon correcte de formuler la tâche.

Visualisation des cartes d'attention

C'est le moment où tout se voit. On récupère les poids d'attention croisée de la dernière couche du décodeur pour une phrase de validation, et on trace une carte de chaleur.

import matplotlib.pyplot as plt

modele.eval()
with torch.no_grad():
src = encoder_source("3 mars 2026")
cbl = generer(modele, src)
_, poids = modele.decodeurs[-1].attention_croisee(...)
# poids : (1, h, L_cbl, L_src) apres passage dans le decodeur

carte = poids[0, 0].numpy() # tete 0
plt.imshow(carte, aspect="auto")
plt.xlabel("jetons source")
plt.ylabel("jetons cible")
plt.colorbar()
plt.title("Attention croisee, tete 0")

La lecture est éloquente : quand le décodeur produit les chiffres de l'année (les positions 0 à 3 de la cible « 2026 »), l'attention pointe vers les jetons « 2 », « 0 », « 2 », « 6 » de la source. Quand il produit le mois (« -03 »), elle pointe vers « mars ». Quand il produit le jour (« -03 »), elle pointe vers « 3 ». Le modèle a appris la structure de la date sans qu'on la lui ait spécifiée, en propageant le signal de gradient à travers l'attention croisée.

Tests unitaires de forme

Un modèle qui ne renvoie pas la bonne forme est un bogue latent : la perte descend, mais on prédit autre chose que ce qu'on croit. Deux ou trois assertions bien placées valent mieux qu'une inspection à posteriori.

def tester_formes():
B, L_src, L_cbl = 4, 12, 10
modele = Transformer(taille_src=30, taille_cbl=15)
src = torch.randint(0, 30, (B, L_src))
cbl = torch.randint(0, 15, (B, L_cbl))

logits = modele(src, cbl)
assert logits.shape == (B, L_cbl, 15), logits.shape

memoire = modele.encoder(src)
assert memoire.shape == (B, L_src, 64), memoire.shape

print("Formes OK")

tester_formes()

Ce test s'ajoute au dépôt à côté du code du modèle. Il se lance en une seconde et il attrape la moitié des régressions futures — un view mal fait, un transpose oublié, une dimension inversée entre encodeur et décodeur.

Un test unitaire de forme dès la première version

Sur un modèle personnalisé, écrivez ce test avant la boucle d'entraînement. Une fois qu'un entraînement tourne sur des formes incorrectes, plus rien n'est fiable, et déboguer prend beaucoup plus de temps que d'écrire les cinq lignes ci-dessus.

Après le fil rouge

Ce petit modèle apprend la tâche jouet en quelques minutes. Pour un projet réel — traduction français-anglais, résumé, classification de tickets — on utilisera un modèle préentraîné via transformers de Hugging Face, quitte à l'affiner sur son jeu de données. Le fil rouge de ce cours n'était pas de construire un modèle de production ; il était de rendre lisible ce qui se passe à l'intérieur de ces modèles, pour que leurs choix d'architecture, leurs pièges d'entraînement et leurs contraintes de service ne soient plus des boîtes noires.

En résumé

  • Un Transformer encodeur-décodeur s'écrit en 200 lignes de PyTorch avec les briques posées aux modules 2 à 5, et s'entraîne sur CPU en quelques minutes sur une tâche jouet.
  • La traduction de dates rend les cartes d'attention lisibles à l'œil nu : chaque jeton cible pointe visiblement vers le fragment source correspondant.
  • Le décalage d'un jeton entre l'entrée et la sortie attendue du décodeur, avec [BOS] et [EOS], est indispensable ; sans lui, la tâche est triviale et le modèle n'apprend rien.
  • Des tests unitaires de forme écrits avant la boucle d'entraînement attrapent la moitié des bogues d'implémentation à venir, en une seconde d'exécution.

Module suivant : le récapitulatif du cours et l'annonce de l'examen final.