Aller au contenu

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

  1. Précision mixte (BF16, et de plus en plus FP8 via Transformer Engine) ;
  2. Fusion des opérations élémentaires (Liger Kernel, torch.compile) ;
  3. Gradient checkpointing : recalculer les activations plutôt que les stocker. Coûte ~30 % de FLOP, économise énormément de mémoire ;
  4. FSDP / ZeRO : partitionner poids, gradients et états d'optimiseur ;
  5. Recouvrement communication/calcul ;
  6. FlashAttention, obligatoire ;
  7. 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 :

\[ t_{\min} = \frac{p\Theta + \text{cache KV}}{B} \]

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 :

  1. surcoût de lancement : 7 à 11 noyaux par couche × 1,3 µs (avec CUDA Graphs) ;
  2. barrières globales implicites : les SM attendent les retardataires du noyau précédent ;
  3. bulles mémoire : au démarrage de chaque noyau, plus rien n'est en vol ;
  4. allers-retours d'activations : sur Llama-1B/B200, 42 % du temps est consacré à stocker, synchroniser et recharger les activations ;
  5. 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 :

  1. 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).
  2. 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

  1. 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.
  2. La GEMM optimisée est une hiérarchie de trois pavages — bloc, thread, instruction. C'est le modèle de toute optimisation dense.
  3. 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.
  4. 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.
  5. La quantification est le levier direct du décodage, parce que le temps est dicté par la lecture des poids.
  6. 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.

\[I = \frac{2\Theta b}{2\Theta + c \cdot b}\]

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