4 · La fusion de noyaux¶
Le levier d'optimisation au meilleur rapport effort/gain de tout le document, et le concept dont les megakernels sont la forme extrême.
4.1 L'arithmétique de la fusion¶
Soit \(n\) opérations élémentaires appliquées successivement à un tenseur de \(m\) éléments de \(s\) octets.
Sans fusion, chaque opération lit et écrit le tenseur entier :
Avec fusion, on lit une fois et on écrit une fois :
Le gain est exactement \(n\).
Exemple concret :
y = torch.relu(x) # lit 4m, écrit 4m
y = y * scale # lit 4m, écrit 4m
y = y + biais # lit 4m, écrit 4m
# total : 24m octets
contre
y = torch.relu(x) * scale + biais # fusionné : 8m octets
Facteur 3. Et pour un modèle avec des dizaines d'opérations élémentaires par couche, le facteur est du même ordre.
4.2 Les quatre types de fusion¶
| Type | Description | Difficulté | Outil |
|---|---|---|---|
| Élémentaire | plusieurs opérations point à point | facile | torch.compile, Triton |
| Réduction + élémentaire | normalisation, softmax | moyenne | Triton |
| GEMM + épilogue | biais, activation, quantification après une GEMM | moyenne | CUTLASS, cuBLASLt |
| GEMM + GEMM | deux produits enchaînés | difficile | FlashAttention, megakernels |
Fusion élémentaire¶
Le cas facile, et le plus rentable. Inductor et Triton le font automatiquement.
Réduction + élémentaire¶
Une RMSNorm :
Naïvement : un noyau pour le carré, un pour la somme, un pour la racine, un pour la division, un pour la multiplication. Cinq allers-retours.
Fusionné : un seul noyau qui charge \(x\) une fois, garde le bloc en registres ou en mémoire partagée, calcule la somme par réduction de warp, et écrit \(y\). Un aller-retour.
@triton.jit
def rmsnorm_kernel(x_ptr, g_ptr, y_ptr, stride, n, eps, BLOC: tl.constexpr):
ligne = tl.program_id(0)
cols = tl.arange(0, BLOC)
masque = cols < n
x = tl.load(x_ptr + ligne * stride + cols, mask=masque, other=0.0)
x = x.to(tl.float32) # accumuler en FP32
var = tl.sum(x * x, axis=0) / n
rstd = 1.0 / tl.sqrt(var + eps)
g = tl.load(g_ptr + cols, mask=masque, other=0.0)
y = x * rstd * g
tl.store(y_ptr + ligne * stride + cols, y.to(tl.float16), mask=masque)
Facteur 5 sur cette opération.
GEMM + épilogue¶
Après une GEMM, on applique souvent un biais, une activation, un résidu, ou une quantification.
Sans fusion : la GEMM écrit \(\mathbf{C}\) en HBM, un second noyau la relit, applique la fonction, réécrit.
Fusionné : l'épilogue s'applique aux accumulateurs encore en registres, avant la seule écriture.
C'est ce que permettent :
cublasLtMatmulavecCUBLASLT_EPILOGUE_GELU_BIASouRELU_BIAS;- CUTLASS avec un
EpilogueOppersonnalisé ; - Triton, en écrivant l'épilogue dans le noyau ;
torch.compile, qui le fait parfois automatiquement.
Le gain : on économise une lecture et une écriture de \(\mathbf{C}\). Pour une GEMM \(4096^3\) en BF16, cela fait \(2 \times 33{,}5\) Mo, soit ~20 µs sur H100. Sur une GEMM qui en dure 140, c'est 14 %.
GEMM + GEMM¶
Le cas difficile, et le plus intéressant.
Fusionner \(\mathbf{C} = \mathbf{A}\mathbf{B}\) puis \(\mathbf{E} = \mathbf{C}\mathbf{D}\) est difficile parce que chaque élément de \(\mathbf{E}\) dépend de toute une ligne de \(\mathbf{C}\), donc de beaucoup de travail de la première GEMM.
FlashAttention est exactement cela : \(\mathbf{Q}\mathbf{K}^\top\) puis \(\operatorname{softmax}(\cdot)\mathbf{V}\), fusionnées grâce à un pavage compatible et au softmax en ligne.
Le MLP à porte est un autre cas :
Les deux premières GEMM partagent \(\mathbf{x}\) : on les concatène en une seule GEMM \((d \times 2d_{\text{ff}})\), ce qui économise une lecture de \(\mathbf{x}\) et un lancement. C'est fait par tous les moteurs d'inférence.
Fusionner la troisième est plus délicat, mais c'est précisément ce que fait un megakernel — en produisant et consommant l'état caché par morceaux, avec un compteur par morceau.
4.3 Quand la fusion n'est pas rentable¶
Les cinq cas où il ne faut pas fusionner
1. Les données sont déjà en cache. Si le tenseur fait 10 Mo et tient en L2 (50 Mo sur H100), les « allers-retours en HBM » n'ont pas lieu : ils sont servis par le L2. Le gain de la fusion s'effondre.
2. La mémoire partagée devient le goulot. Fusionner impose de garder les intermédiaires quelque part. Si cela fait passer de 4 blocs résidents à 1, on perd plus qu'on ne gagne.
3. La pression sur les registres augmente. Un épilogue complexe peut provoquer du spilling, ce qui annule tout.
4. L'opération fusionnée est limitée par le calcul. Fusionner deux grosses GEMM n'apporte rien : elles sont déjà limitées par le calcul, et le trafic de \(\mathbf{C}\) est marginal.
5. La forme empêche le pavage compatible. Si l'opération 2 a besoin d'une réduction globale sur le résultat de l'opération 1, la fusion exige une synchronisation inter-blocs — donc un megakernel.
4.4 Ce que fait Inductor, et ce qu'il ne fait pas¶
torch.compile fusionne automatiquement :
| Motif | Fusionné ? |
|---|---|
| Chaîne d'opérations élémentaires | oui |
| Élémentaire + réduction | oui, souvent |
| Réduction + élémentaire | oui, souvent |
| Deux réductions successives | parfois |
| GEMM + épilogue élémentaire | parfois (dépend de la version et du motif) |
| GEMM + GEMM | non |
| Attention complète | non (il appelle un noyau existant) |
| À travers une frontière de graphe | non |
Pour vérifier ce qui a été fusionné :
import os
os.environ["TORCH_LOGS"] = "output_code"
modele_c = torch.compile(modele)
modele_c(entree) # affiche le code Triton généré
Le nombre de fonctions triton_poi_fused_* dans la sortie vous dit combien de
noyaux subsistent, et leurs noms indiquent ce qui a été regroupé.
4.5 Les bibliothèques de noyaux fusionnés¶
Plutôt que d'écrire les vôtres :
| Bibliothèque | Contenu |
|---|---|
| Liger Kernel (LinkedIn) | RMSNorm, RoPE, SwiGLU, GeGLU, entropie croisée fusionnée, LayerNorm — pour l'entraînement de LLM |
| FlashInfer | attention, échantillonnage, quantification, décodage |
| xFormers | attention et blocs de transformeur |
| Apex (NVIDIA) | LayerNorm fusionnée, optimiseurs fusionnés |
| Unsloth | noyaux d'affinage économes en mémoire |
| Transformer Engine | FP8 de bout en bout, avec gestion des échelles |
L'entropie croisée fusionnée de Liger mérite un mot : dans un LLM à grand vocabulaire, la matrice des logits \((b \times S \times V)\) est énorme — pour \(V = 128\,000\), \(b S = 8192\), en FP32 : 4,2 Go. La fusionner avec le calcul de la perte évite de la matérialiser, ce qui économie plusieurs gigaoctets et autant de trafic.
C'est le même principe que FlashAttention, appliqué à la dernière couche.
4.6 De la fusion au megakernel¶
Poussons la logique.
Niveau 0 : un noyau par opération ← PyTorch impératif
Niveau 1 : fusion élémentaire ← torch.compile
Niveau 2 : fusion réduction + élémentaire ← Triton, Liger
Niveau 3 : fusion GEMM + épilogue ← CUTLASS, cuBLASLt
Niveau 4 : fusion GEMM + GEMM ← FlashAttention
Niveau 5 : fusion d'une couche entière ← ?
Niveau 6 : fusion du MODÈLE ENTIER ← MEGAKERNEL
Chaque niveau supprime des allers-retours en HBM. Le niveau 6 les supprime tous, sauf ceux qui sont physiquement nécessaires (lire les poids, écrire le résultat).
Ce qui bloque au niveau 5
Le passage du niveau 4 au niveau 5 exige de franchir une barrière : la synchronisation inter-blocs.
Tant qu'on reste dans un noyau, la fusion est limitée à ce qu'un bloc peut calculer seul. Une couche de transformeur ne le peut pas : la projection de sortie a besoin du résultat de l'attention de toutes les têtes, calculées par des blocs différents.
Il faut donc soit une frontière de noyau (le coût qu'on cherche à éviter), soit un mécanisme de dépendance fine — c'est-à-dire un megakernel.
C'est exactement le raisonnement qui a mené aux travaux de la partie 8.
Résumé du chapitre¶
À retenir
- Fusionner \(n\) opérations élémentaires divise le trafic par \(n\). C'est le levier au meilleur rapport effort/gain.
- Quatre types : élémentaire (facile), réduction + élémentaire (Triton), GEMM + épilogue (CUTLASS/cuBLASLt), GEMM + GEMM (difficile).
- FlashAttention est une fusion GEMM + GEMM, rendue possible par le softmax en ligne.
- Ne pas fusionner quand : les données tiennent en L2, la mémoire partagée ou les registres deviennent limitants, l'opération est déjà limitée par le calcul, ou le pavage est incompatible.
- Inductor fusionne l'élémentaire et une partie des réductions, pas deux GEMM ni l'attention.
- Utilisez Liger Kernel, FlashInfer, xFormers, Transformer Engine avant d'écrire les vôtres.
- Le megakernel est le niveau 6 de la fusion : le modèle entier. Ce qui bloque avant, c'est la synchronisation inter-blocs.
Vérifiez que vous avez compris¶
Vous fusionnez cinq opérations élémentaires et gagnez seulement 1,3× au lieu de 5×. Pourquoi ?
Trois causes plausibles, à vérifier dans l'ordre :
- Les données tenaient en L2. Avec un tenseur de 8 Mo sur un H100 (50 Mo de L2), les cinq opérations lisaient déjà depuis le L2, pas depuis la HBM. Le trafic HBM était de 2 allers-retours, pas 10.
- Le noyau n'était pas limité par la mémoire — par exemple s'il
contient une
expsur la SFU, qui domine. - Le surcoût de lancement dominait. Cinq noyaux de 3 µs chacun coûtent 15 µs de lancement pour 15 µs de calcul ; les fusionner supprime le lancement mais pas le calcul.
Mesurez dram__bytes.sum avant et après : si le trafic HBM n'a pas été
divisé par 5, c'est la cause 1.
Pourquoi fusionner deux GEMM 4096³ n'apporte-t-il presque rien ?
Parce qu'elles sont limitées par le calcul.
Chaque GEMM fait \(1{,}37\times10^{11}\) FLOP pour ~\(2\times10^8\) octets. Le temps est dicté par les FLOP (~140 µs en BF16), et le trafic de la matrice intermédiaire \(\mathbf{C}\) (33,5 Mo, soit ~10 µs de lecture + 10 µs d'écriture) représente ~7 % du total.
Fusionner économiserait donc au mieux 7 %, pour une complexité considérable (il faut un pavage compatible entre les deux GEMM, ce qui contraint fortement les tailles de tuile et dégrade souvent la première GEMM).
La fusion est rentable quand on est limité par la mémoire. Toujours revenir au roofline.
L'entropie croisée fusionnée de Liger économise plusieurs gigaoctets. Comment, alors que les logits sont nécessaires au calcul du gradient ?
Parce qu'elle traite le vocabulaire par morceaux.
Au lieu de calculer tous les logits \((bS \times V)\), puis le softmax, puis la perte, puis le gradient, elle boucle sur des tranches du vocabulaire :
- calculer les logits d'une tranche \((bS \times V_{\text{chunk}})\) ;
- accumuler les statistiques du log-sum-exp (comme le softmax en ligne de FlashAttention) ;
- calculer immédiatement la contribution au gradient de cette tranche ;
- jeter les logits de la tranche.
Seules les statistiques \(O(bS)\) et le gradient d'entrée sont conservés. C'est le même principe que FlashAttention — pavage plus statistiques en ligne — appliqué à la dernière couche.
Chapitre suivant : 5 · La quantification