8 · Entraînement contre inférence¶
Deux régimes structurellement différents, souvent confondus. Ce chapitre les sépare proprement et conclut la partie en établissant pourquoi le décodage à petit lot mène directement aux megakernels.
8.1 Le tableau des différences¶
| Entraînement | Préremplissage | Décodage | |
|---|---|---|---|
| Jetons traités par passe | \(bS\) | \(bS\) | \(b\) |
| FLOP | \(6\Theta bS\) | \(2\Theta bS\) | \(2\Theta b\) |
| Octets (poids) | \(2\Theta\) | \(2\Theta\) | \(2\Theta\) |
| Intensité arithmétique | \(\approx 3bS\) | \(\approx bS\) | \(\approx b\) |
| Régime | calcul | calcul | mémoire |
| Mémoire dominée par | activations + états d'optimiseur | activations | poids + cache KV |
| Métrique | jetons/s, MFU | temps au premier jeton | jetons/s par requête |
| Sensible au surcoût de lancement | non | non | oui |
C'est la ligne « intensité arithmétique » qui structure tout.
8.2 L'entraînement¶
Le régime¶
\(I \approx 3bS\). Avec \(b = 8\) et \(S = 4096\) : \(I \approx 98\,000\), très au-dessus du seuil de 296. Massivement limité par le calcul.
Les conséquences :
- les tensor cores sont la ressource critique ;
- la métrique de référence est la MFU (Model FLOPs Utilization) : la fraction du pic matériel effectivement utilisée ;
- un entraînement bien optimisé atteint 40 à 55 % de MFU. Au-delà de 60 %, c'est exceptionnel.
Où passe le temps¶
| Poste | Part typique |
|---|---|
| GEMM (avant + arrière) | 60-75 % |
| Attention | 10-25 % (croît avec \(S\)) |
| Communication (all-reduce, all-gather) | 5-20 % |
| Opérations élémentaires, normalisations | 5-10 % |
| Optimiseur | 2-5 % |
Les leviers¶
- Précision mixte (BF16, et de plus en plus FP8 via Transformer Engine) ;
- Fusion des opérations élémentaires (Liger Kernel,
torch.compile) ; - Gradient checkpointing : recalculer les activations plutôt que les stocker. Coûte ~30 % de FLOP, économise énormément de mémoire ;
- FSDP / ZeRO : partitionner poids, gradients et états d'optimiseur ;
- Recouvrement communication/calcul ;
- FlashAttention, obligatoire ;
- Chargement des données : DALI, préchargement — souvent le vrai goulot.
Le goulot le plus fréquent en entraînement n'est pas le modèle
Sur un entraînement de vision ou de multimodal, le chargement et l'augmentation des données sur CPU sont fréquemment limitants. Le symptôme : des trous réguliers dans la chronologie Nsight Systems.
Vérifiez toujours l'utilisation GPU avant d'optimiser un noyau.
8.3 Le préremplissage¶
Le traitement de l'invite, avant la génération.
\(I \approx bS\). Avec un seul utilisateur et une invite de 2 000 jetons : \(I = 2000 \gg 296\). Limité par le calcul, même à lot 1.
Les conséquences :
- c'est presque le même problème que l'entraînement (passe avant seulement) ;
- la métrique est le temps au premier jeton (TTFT) ;
- l'attention y coûte \(O(S^2)\), donc FlashAttention est déterminant sur les longues invites ;
- le préremplissage saturé peut bloquer le décodage : d'où la désagrégation (prefill/decode disaggregation), qui consiste à les exécuter sur des GPU différents.
8.4 Le décodage : le régime qui compte¶
\(I \approx b\). C'est là que tout se joue.
Le plancher physique¶
Pour un modèle de \(\Theta\) paramètres en précision \(p\) octets, à lot 1 :
Sur H100 (3,35 To/s), en BF16 :
| Modèle | Poids | \(t_{\min}\) | Jetons/s max |
|---|---|---|---|
| 1 B | 2 Go | 0,60 ms | 1 675 |
| 8 B | 16 Go | 4,78 ms | 209 |
| 70 B (sur 2 GPU) | 70 Go/GPU | 20,9 ms | 48 |
| 405 B (sur 8 GPU) | 101 Go/GPU | 30,2 ms | 33 |
Ce sont des bornes supérieures physiques. Aucune implémentation ne peut les dépasser sans réduire les octets lus (quantification, MoE, cache compressé).
L'écart entre le plancher et la réalité¶
C'est le chiffre central de tout ce document.
| Système | Fraction de la bande passante atteinte |
|---|---|
| PyTorch impératif | 20-35 % |
| vLLM / SGLang | ~50 % |
| Megakernel (Hazy Research, H100) | 78 % |
Les 50 % perdus par les systèmes classiques se décomposent en :
- surcoût de lancement : 7 à 11 noyaux par couche × 1,3 µs (avec CUDA Graphs) ;
- barrières globales implicites : les SM attendent les retardataires du noyau précédent ;
- bulles mémoire : au démarrage de chaque noyau, plus rien n'est en vol ;
- allers-retours d'activations : sur Llama-1B/B200, 42 % du temps est consacré à stocker, synchroniser et recharger les activations ;
- synchronisation entre warps : ~60 ns par barrière, 40 µs sur 600.
Le raisonnement complet, en une image
Bande passante H100 : 3 350 Go/s
│
├─ 100 % → 1 350 passes avant/s sur Llama-1B (plancher physique)
│
├─ 50 % → ~770 passes/s ← vLLM, SGLang
│ perdu en : lancements, barrières, bulles, activations
│
└─ 78 % → ~1 050 passes/s ← MEGAKERNEL
soit < 1 ms par passe avant sur H100
Il n'y a rien d'autre à gagner : le megakernel ne calcule pas mieux, il supprime le temps où rien ne circule.
8.5 Le regroupement, et sa limite¶
Grouper les requêtes augmente \(I\) linéairement :
| Lot | \(I\) | Régime (H100, BF16) |
|---|---|---|
| 1 | 1 | mémoire (×296 sous le seuil) |
| 16 | 16 | mémoire (×18) |
| 64 | 64 | mémoire (×4,6) |
| 296 | 296 | seuil |
| 512 | 512 | calcul |
C'est le fondement du continuous batching de vLLM et SGLang : accumuler autant de requêtes concurrentes que possible.
Mais deux limites :
- Le cache KV. Il croît avec le lot et avec le contexte. À contexte long, \(I\) plafonne (voir le calcul du chapitre 1, §1.4).
- La latence par requête. Un grand lot améliore le débit total et dégrade le temps entre jetons pour chaque requête. Pour un usage interactif, c'est inacceptable.
D'où la distinction fondamentale de la partie 8 :
| Objectif | Contrainte | Megakernel adapté |
|---|---|---|
| Latence (lot 1-8) | temps par jeton | interpréteur latence, MPK |
| Débit (lot 1000+) | jetons/s total | megakernel TP |
Les deux existent, et ce sont des objets techniques différents.
8.6 Le décodage spéculatif¶
Une technique qui change l'arithmétique et mérite d'être connue ici.
L'idée : un petit modèle « brouillon » propose \(k\) jetons ; le grand modèle les vérifie en une seule passe, puis on accepte le plus long préfixe correct.
L'arithmétique :
- sans spéculation : \(k\) passes du grand modèle pour \(k\) jetons ;
- avec spéculation : 1 passe du grand modèle (avec \(k\) jetons en entrée) + \(k\) passes du petit modèle.
La passe de vérification traite \(k\) jetons, donc \(I \approx k\) au lieu de 1. On convertit un problème limité par la mémoire en un problème un peu plus équilibré, en utilisant des tensor cores qui étaient de toute façon inactifs.
Le gain effectif dépend du taux d'acceptation. Avec \(k = 5\) et 70 % d'acceptation, on génère ~3,5 jetons par passe du grand modèle : accélération de ~2,5×.
Les variantes : Medusa (têtes supplémentaires sur le même modèle), EAGLE (prédiction au niveau des caractéristiques), prédiction multi-jetons (intégrée à l'architecture, comme dans DeepSeek-V3 et Qwen3).
L'interaction avec les megakernels
Le décodage spéculatif et les megakernels sont complémentaires : le premier augmente le travail utile par passe, le second supprime le surcoût de chaque passe.
Ils sont aussi en tension : la spéculation introduit du dynamisme (le nombre de jetons acceptés varie), ce qui complique la compilation statique d'un megakernel. C'est précisément le problème que le papier Event Tensor attaque.
8.7 Le tableau de synthèse des optimisations¶
| Optimisation | Entraînement | Préremplissage | Décodage |
|---|---|---|---|
| Tensor cores | essentiel | essentiel | peu utile à lot 1 |
| FlashAttention | essentiel | essentiel | Flash-Decoding |
| Fusion élémentaire | important | important | essentiel |
| Quantification des poids | limitée (FP8) | utile | essentiel |
| Quantification du cache KV | — | — | essentiel en contexte long |
| MoE | utile | utile | essentiel |
| Gradient checkpointing | essentiel | — | — |
| Regroupement | naturel | limité | essentiel pour le débit |
| CUDA Graphs | utile | peu utile | essentiel |
| Megakernel | non pertinent | peu utile | le levier restant |
| Décodage spéculatif | — | — | très utile |
Résumé de la partie 7¶
Les six idées à emporter
- Tout se joue sur l'intensité arithmétique. Entraînement et préremplissage : \(I\) énorme, limité par le calcul. Décodage à lot 1 : \(I \approx 1\), limité par la mémoire d'un facteur ~296.
- La GEMM optimisée est une hiérarchie de trois pavages — bloc, thread, instruction. C'est le modèle de toute optimisation dense.
- FlashAttention ne matérialise jamais \(\mathbf{S}\), grâce au softmax
en ligne qui est exact. FA4 pousse la co-conception jusqu'à déplacer
exp()de la SFU vers les FMA. - La fusion divise le trafic par le nombre d'opérations fusionnées. C'est le meilleur rapport effort/gain, et le megakernel en est la forme extrême.
- La quantification est le levier direct du décodage, parce que le temps est dicté par la lecture des poids.
- Les systèmes actuels atteignent ~50 % de la bande passante en décodage ; le plancher physique est 100 %, et un megakernel atteint 78 %. C'est cet écart qui justifie la partie 8.
Vérifiez que vous avez compris¶
Pourquoi la MFU d'un entraînement plafonne-t-elle vers 50 % même bien optimisé ?
Parce que la MFU compte les FLOP « utiles » du modèle rapportés au pic matériel, et que plusieurs postes ne sont pas comptés comme utiles :
- l'attention, dont le coût réel dépasse souvent le décompte théorique (masquage, softmax, opérations non-matmul) ;
- les opérations élémentaires et les normalisations, limitées par la mémoire, où l'utilisation est de quelques pourcents ;
- la communication, si elle n'est pas entièrement recouverte ;
- le gradient checkpointing, qui ajoute ~30 % de FLOP non comptés ;
- la quantification de vagues et les tailles de tuiles non idéales ;
- le chargement des données.
Atteindre 50 % de MFU signifie donc que les GEMM tournent à 70-85 % du pic, ce qui est très bon.
À contexte 128 k, augmenter le lot n'améliore plus le débit. Pourquoi ?
Parce que le cache KV croît proportionnellement au lot, exactement comme le calcul.
où \(c\) est la taille du cache par séquence. Quand \(c \cdot b \gg 2\Theta\), l'expression tend vers \(2\Theta/c\) — une constante indépendante de \(b\).
Pour un modèle de 8 milliards de paramètres à 128 k de contexte : \(2\Theta = 16\) Go, \(c = 17{,}2\) Go. Dès \(b = 1\), le cache égale les poids ; dès \(b = 4\), il les domine, et \(I\) plafonne à ~\(16/17{,}2 \approx 0{,}9\) par jeton de lot.
Les remèdes sont ceux du chapitre 1 : GQA/MQA, MLA, quantification du cache, attention creuse, ou architectures hybrides.
Un megakernel améliore-t-il l'entraînement ?
Marginalement, et ce n'est pas son objet.
L'entraînement est limité par le calcul : \(I \approx 3bS\), soit des dizaines de milliers. Les tensor cores sont saturés, et le surcoût de lancement (quelques centaines de microsecondes) est négligeable devant un pas d'entraînement qui dure des centaines de millisecondes.
Ce qui limite l'entraînement, c'est la qualité des GEMM (déjà excellente avec cuBLAS/CUTLASS), le recouvrement de la communication, et le chargement des données.
Les megakernels ciblent le régime limité par la mémoire à faible travail par passe — c'est-à-dire le décodage. C'est aussi pourquoi tous les travaux de la partie 8 évaluent sur du décodage, jamais sur de l'entraînement.
Partie suivante : Megakernels
Sources de ce chapitre¶
- Hazy Research, Look Ma, No Bubbles! — 78 % contre ~50 %, le décompte des 600 µs.
- Hazy Research, We Bought the Whole GPU
- Leviathan, Kalman, Matias, Fast Inference from Transformers via Speculative Decoding — arXiv:2211.17192
- Cai et al., Medusa — arXiv:2401.10774
- Kwon et al., Efficient Memory Management for LLM Serving with PagedAttention (vLLM) — arXiv:2309.06180
- Memory-Bound but Not Bandwidth-Limited: The Physical AI Inference Gap in Batch-1 LLM Decode, arXiv:2605.30571