Aller au contenu

1 · Anatomie d'un transformeur sur GPU

Où passe le temps, opération par opération. Ce chapitre construit le décompte qui sert de référence à toute la partie.


1.1 Le squelette

Une couche de transformeur, dans sa forme moderne (pré-normalisation, GLU) :

entrée x  (b, S, d)
   │
   ├─ RMSNorm ────────────────────┐
   │                              │
   ├─ Wq, Wk, Wv  ──→ Q, K, V     │   projections
   ├─ RoPE (rotation positionnelle)│
   ├─ Attention(Q, K, V)          │
   ├─ Wo  ──────────────────────→ │   projection de sortie
   │                              │
   └─ + résidu ←──────────────────┘
   │
   ├─ RMSNorm ────────────────────┐
   │                              │
   ├─ Wgate, Wup ──→ SwiGLU       │   MLP
   ├─ Wdown                       │
   │                              │
   └─ + résidu ←──────────────────┘
   │
sortie  (b, S, d)

Notation utilisée dans tout ce chapitre :

Symbole Signification
\(b\) taille du lot
\(S\) longueur de séquence
\(d\) dimension du modèle (hidden size)
\(h\) nombre de têtes d'attention
\(d_h = d/h\) dimension par tête
\(d_{\text{ff}}\) dimension du MLP, typiquement \(\approx 3{,}5d\) avec SwiGLU
\(L\) nombre de couches
\(\Theta\) nombre total de paramètres

1.2 Le décompte des opérations

Les projections

Chaque projection est une GEMM \((bS \times d) \cdot (d \times d')\) :

Projection Dimensions FLOP
\(\mathbf{W}_q\) \(d \times d\) \(2bSd^2\)
\(\mathbf{W}_k\), \(\mathbf{W}_v\) \(d \times d_{kv}\) \(2bSd \cdot d_{kv}\) chacune
\(\mathbf{W}_o\) \(d \times d\) \(2bSd^2\)
\(\mathbf{W}_{\text{gate}}\), \(\mathbf{W}_{\text{up}}\) \(d \times d_{\text{ff}}\) \(2bSd\,d_{\text{ff}}\) chacune
\(\mathbf{W}_{\text{down}}\) \(d_{\text{ff}} \times d\) \(2bSd\,d_{\text{ff}}\)

Avec l'attention à requêtes groupées (GQA), \(d_{kv} = d \cdot g/h\) où \(g\) est le nombre de groupes de clés-valeurs — typiquement 8 pour \(h = 64\), donc \(d_{kv} = d/8\).

L'attention

\[ \operatorname{Attention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \operatorname{softmax}\!\left( \frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_h}} \right) \mathbf{V} \]

Deux GEMM par tête :

  • \(\mathbf{Q}\mathbf{K}^\top\) : \((S \times d_h) \cdot (d_h \times S)\) → \(2 S^2 d_h\) FLOP ;
  • \(\operatorname{softmax}(\cdot)\mathbf{V}\) : \((S \times S) \cdot (S \times d_h)\) → \(2 S^2 d_h\) FLOP.

Sur \(h\) têtes et \(b\) séquences : \(4 b h S^2 d_h = 4 b S^2 d\) FLOP.

Avec un masque causal, la moitié est inutile : \(2bS^2d\) en pratique.

La croissance quadratique, en pratique

L'attention coûte \(O(S^2 d)\) tandis que les projections coûtent \(O(S d^2)\).

Le point de bascule est en \(S \approx d\). Pour un modèle avec \(d = 4096\) :

  • \(S = 512\) : l'attention représente ~10 % du calcul ;
  • \(S = 4\,096\) : ~50 % ;
  • \(S = 32\,768\) : ~85 % ;
  • \(S = 131\,072\) : ~96 %.

C'est pourquoi le contexte long est un problème d'attention, et pourquoi FlashAttention a été si important.

Le total

En sommant sur une couche, avec \(d_{\text{ff}} = 3{,}5d\) et GQA à 8 groupes :

\[ \text{FLOP}_{\text{couche}} \approx \underbrace{2bSd^2 \left(2 + \tfrac{2}{8}\right)}_{\text{attn. proj.}} + \underbrace{2bS^2 d}_{\text{attention}} + \underbrace{3 \times 2bSd \cdot 3{,}5d}_{\text{MLP}} \]
\[ \approx 2bSd^2 (2 + 0{,}25 + 10{,}5) + 2bS^2 d = 25{,}5\, bSd^2 + 2bS^2 d \]

Le MLP domine largement à contexte court : 10,5 des 12,75 unités de \(bSd^2\), soit 82 % du calcul des GEMM.

La règle des 6N

Pour l'entraînement, une approximation universellement utilisée :

\[\text{FLOP} \approx 6 \, \Theta \, N_{\text{jetons}}\]

Le facteur 6 = 2 (passe avant) + 4 (passe arrière, environ deux fois le coût de la passe avant). Elle néglige l'attention, ce qui est valide tant que \(S \ll d\).

Pour l'inférence en préremplissage : \(\approx 2\Theta N_{\text{jetons}}\).


1.3 Le décompte des octets

C'est le décompte qui compte vraiment, parce que c'est lui qui décide du régime.

En entraînement

Catégorie Volume (BF16)
Poids \(2\Theta\)
Gradients \(2\Theta\)
États d'optimiseur (Adam) \(8\Theta\) (deux moments en FP32)
Copie maîtresse FP32 \(4\Theta\)
Total statique \(\approx 16\Theta\)
Activations \(\propto b S d L\)

Pour un modèle de 8 milliards de paramètres : 128 Go rien qu'en état statique, avant les activations. D'où le parallélisme de données à état partitionné (ZeRO/FSDP) et le gradient checkpointing.

En inférence, préremplissage

On lit les poids une fois pour traiter \(bS\) jetons :

\[ I \approx \frac{2 \Theta \cdot bS}{2\Theta} = bS \]

Avec \(bS = 4\,096\) jetons, \(I = 4\,096 \gg 296\) : limité par le calcul.

En inférence, décodage

On lit les poids une fois pour produire \(b\) jetons :

\[ I \approx \frac{2\Theta b}{2\Theta + \text{cache KV}} \approx b \]

Avec \(b = 1\) : \(I \approx 1\). Limité par la mémoire, d'un facteur ~296.


1.4 Le cache clé-valeur

Souvent sous-estimé, il devient dominant à contexte long.

\[ \text{taille}_{\text{cache}} = 2 \times b \times S \times L \times d_{kv} \times \text{octets/élément} \]

Le facteur 2 compte les clés et les valeurs.

Exemple concret — modèle de 8 milliards de paramètres, \(L = 32\), \(d = 4096\), GQA à 8 groupes sur 32 têtes (donc \(d_{kv} = 1024\)), en BF16 :

Contexte Cache par séquence
4 096 0,54 Go
32 768 4,3 Go
131 072 17,2 Go
1 048 576 137 Go

Les poids font 16 Go. À partir de ~128 000 jetons, le cache dépasse les poids.

Et il est lu intégralement à chaque jeton généré, ce qui l'ajoute au dénominateur de l'intensité arithmétique :

\[ I_{\text{décodage}} \approx \frac{2\Theta b}{2\Theta + \text{cache}} \]

La conséquence pratique

À contexte long, augmenter le lot n'améliore plus l'intensité arithmétique, parce que le cache croît avec le lot au même rythme que le calcul. On reste bloqué en régime limité par la mémoire quoi qu'on fasse.

C'est ce qui motive : GQA et MQA (réduire \(d_{kv}\)), MLA de DeepSeek (compresser le cache), la quantification du cache (FP8, INT4), l'attention creuse, et les architectures hybrides mêlant attention et récurrence.


1.5 Où passe réellement le temps

Le meilleur décompte publié est celui de Hazy Research pour Llama-1B sur B200, passe avant complète en 600 µs :

Poste Temps Part
Stockage des activations, attente de cohérence, rechargement 250 µs 42 %
RMSNorm et produits matrice-vecteur (95 % pour les matvec) 200 µs 33 %
Attente du chargement des poids depuis la mémoire globale 30 µs 5 %
Synchronisation bas niveau entre warps (~60 ns par barrière) 40 µs 7 %
Configuration et divers 80 µs 13 %

Le chiffre le plus important de ce chapitre

42 % du temps est passé à écrire des activations, attendre leur cohérence, et les relire.

Ce n'est pas du calcul. Ce n'est même pas de la lecture de poids. C'est le coût pur des frontières entre opérations.

Sur cette base, la conclusion est directe : si l'on pouvait garder les activations en mémoire partagée ou en registres entre deux opérations, on supprimerait la moitié du temps d'exécution.

C'est exactement ce que fait un megakernel.

L'analyse complémentaire : la bande passante d'un H100 permettrait environ 1 350 passes avant par seconde sur Llama-1B, alors que les systèmes réels en réalisent ~770 — avec « environ cinq microsecondes de blocage par noyau, sur 7 lancements par couche et 16 couches ».


1.6 La liste des noyaux d'une passe avant

Pour situer, voici ce qu'exécute une implémentation classique par couche :

# Noyau Régime
1 RMSNorm mémoire
2 GEMM QKV (souvent fusionnée) calcul (prefill) / mémoire (décodage)
3 RoPE mémoire
4 Attention (FlashAttention ou FlashInfer) mixte
5 GEMM projection de sortie idem 2
6 Résidu (souvent fusionné avec 7) mémoire
7 RMSNorm mémoire
8 GEMM gate + up (fusionnée) idem 2
9 SwiGLU mémoire
10 GEMM down idem 2
11 Résidu mémoire

Soit 7 à 11 lancements par couche, et 224 à 352 pour un modèle de 32 couches. À 1,3 µs par lancement avec CUDA Graphs : 291 à 458 µs de pur surcoût de lancement.

Pour un modèle où le décodage d'un jeton devrait prendre ~5 ms, cela représente 6 à 9 %. Pour un petit modèle où il devrait prendre 1 ms, c'est 30 à 45 %.


Résumé du chapitre

À retenir

  • Le MLP domine le calcul à contexte court : ~82 % des FLOP de GEMM. L'attention devient dominante au-delà de \(S \approx d\).
  • Règle des 6N pour l'entraînement, 2N pour le préremplissage.
  • L'état d'entraînement fait ~16 octets par paramètre (poids, gradients, moments Adam, copie maîtresse).
  • En décodage à lot 1, \(I \approx 1\) contre un seuil de ~296 : 296 fois du mauvais côté.
  • Le cache KV dépasse les poids au-delà de ~128 k jetons pour un modèle de 8 milliards de paramètres, et il annule le bénéfice du regroupement.
  • Sur B200, 42 % du temps d'une passe avant de Llama-1B est consacré à stocker, synchroniser et recharger les activations — le coût des frontières entre opérations.
  • 7 à 11 lancements de noyau par couche, soit 291 à 458 µs de surcoût pour 32 couches même avec CUDA Graphs.

Vérifiez que vous avez compris

Pourquoi GQA réduit-elle plus le cache KV que le calcul ?

Parce que GQA partage les têtes de clés et de valeurs entre plusieurs têtes de requêtes, sans réduire le nombre de têtes de requêtes.

  • Cache : proportionnel à \(d_{kv}\), divisé par \(h/g\). Avec 32 têtes et 8 groupes : divisé par 4.
  • Calcul de l'attention : chaque tête de requête calcule toujours son attention complète. Le nombre de FLOP est inchangé.
  • Calcul des projections K et V : divisé par 4, mais elles ne représentent qu'une petite partie du total.

GQA est donc une optimisation mémoire, ce qui est exactement ce dont l'inférence a besoin. C'est pour cela qu'elle est universelle depuis Llama 2.

Un modèle de 70 milliards de paramètres en BF16, décodage à lot 32, contexte 8 192. Limité par quoi sur un H100 ?
  • Poids : \(2 \times 70 \times 10^9 = 140\) Go — cela ne tient pas sur un H100 (80 Go), il faut au moins 2 GPU. Prenons le cas d'un parallélisme de tenseurs sur 2 cartes, soit 70 Go de poids par carte.
  • Cache KV : avec \(L = 80\), \(d_{kv} = 1024\) : \(2 \times 32 \times 8192 \times 80 \times 1024 \times 2 = 86\) Go au total, 43 Go par carte.
  • Octets lus par pas de décodage, par carte : \(70 + 43 = 113\) Go.
  • FLOP par pas, par carte : \(2 \times 35 \times 10^9 \times 32 = 2{,}24\) TFLOP.
  • \(I = 2{,}24 \times 10^{12} / (113 \times 10^9) \approx 20\).

Contre \(I_{\text{crit}} = 296\) : toujours limité par la mémoire, d'un facteur 15, malgré un lot de 32. Le cache KV est ici la moitié du trafic.

C'est exactement l'illustration du §1.4 : à contexte long, le regroupement ne suffit plus.

Les 42 % de temps passés à gérer les activations sur B200 : pourquoi B200 et pas H100 ?

Parce que le problème s'aggrave avec les générations récentes.

Sur B200, le débit des tensor cores a plus que doublé par rapport à H100, tandis que la bande passante mémoire a été multipliée par ~2,4 et que les latences de synchronisation n'ont pas fondamentalement changé.

Résultat : la partie « calcul utile » d'une passe avant raccourcit, la partie « frontières et synchronisation » beaucoup moins. Sa part relative augmente.

C'est la raison structurelle pour laquelle les megakernels sont apparus en 2025 et pas en 2020 : le surcoût était toujours là, mais il était noyé dans un calcul plus lent.


Chapitre suivant : 2 · La GEMM, de zéro à cuBLAS


Sources de ce chapitre