Aller au contenu

2 · La GEMM, de zéro à cuBLAS

L'exercice canonique de la programmation GPU. Huit étapes, de 1 % à plus de 90 % du pic, chacune enseignant un principe qui se réapplique partout.


2.1 Le problème et sa borne

\[ \mathbf{C}_{M \times N} = \alpha \, \mathbf{A}_{M \times K} \mathbf{B}_{K \times N} + \beta \, \mathbf{C} \]

Prenons \(M = N = K = 4096\) en FP32 sur un H100.

Travail : \(2MNK = 1{,}37 \times 10^{11}\) FLOP.

Octets minimaux : \(3 \times 4096^2 \times 4 = 2{,}01 \times 10^8\).

Intensité arithmétique maximale : \(I = 683\), très au-dessus du seuil de 20. La GEMM est le problème compute-bound par excellence — à condition d'être bien écrite.

Borne de temps : \(1{,}37\times10^{11} / (67 \times 10^{12}) = 2{,}05\) ms en FP32 vectoriel, ou \(138\ \mu s\) en BF16 tensor.


2.2 Étape 0 — la version naïve

__global__ void sgemm_naif(int M, int N, int K, float alpha,
                           const float* A, const float* B,
                           float beta, float* C) {
    int x = blockIdx.x * blockDim.x + threadIdx.x;   // colonne
    int y = blockIdx.y * blockDim.y + threadIdx.y;   // ligne
    if (x >= N || y >= M) return;

    float acc = 0.0f;
    for (int k = 0; k < K; ++k)
        acc += A[y*K + k] * B[k*N + x];

    C[y*N + x] = alpha * acc + beta * C[y*N + x];
}

Intensité arithmétique : chaque thread fait \(2K\) opérations et lit \(8K\) octets. \(I = 1/4\).

Performance attendue : \(0{,}25 \times 3{,}35\) To/s \(= 0{,}84\) TFLOPS, soit 1,3 % du pic.

Le problème : chaque élément de \(\mathbf{A}\) est lu \(N\) fois, chaque élément de \(\mathbf{B}\) est lu \(M\) fois.


2.3 Étape 1 — la coalescence

Le noyau naïf a un défaut plus grave encore. Avec x = threadIdx.x indexant les colonnes, l'accès A[y*K + k] a y constant sur le warp — tous les threads lisent le même élément (diffusion, c'est bien) — et B[k*N + x] a x variable — accès contigus, c'est bien aussi.

Mais si l'on inverse (x indexant les lignes), les deux accès deviennent catastrophiques. C'est l'erreur du chapitre 1 de la partie 2.

Gain d'une indexation correcte : facteur ~8. On passe à ~10 % du pic FP32 vectoriel. C'est l'optimisation au meilleur rapport effort/gain de toute la séquence : elle consiste à échanger deux lignes.


2.4 Étape 2 — le pavage en mémoire partagée

#define T 32

__global__ void sgemm_tuiles(int M, int N, int K, float alpha,
                             const float* A, const float* B,
                             float beta, float* C) {
    __shared__ float As[T][T];
    __shared__ float Bs[T][T];

    int tx = threadIdx.x, ty = threadIdx.y;
    int lig = blockIdx.y * T + ty;
    int col = blockIdx.x * T + tx;

    float acc = 0.0f;
    for (int t = 0; t < K; t += T) {
        As[ty][tx] = (lig < M && t+tx < K) ? A[lig*K + t+tx] : 0.0f;
        Bs[ty][tx] = (t+ty < K && col < N) ? B[(t+ty)*N + col] : 0.0f;
        __syncthreads();

        #pragma unroll
        for (int k = 0; k < T; ++k)
            acc += As[ty][k] * Bs[k][tx];
        __syncthreads();
    }
    if (lig < M && col < N)
        C[lig*N + col] = alpha * acc + beta * C[lig*N + col];
}

Nouvelle intensité arithmétique : \(I = T/4 = 8\).

Gain : facteur 32 sur le trafic HBM. On atteint typiquement 25 à 35 % du pic FP32 vectoriel.

Le nouveau goulot : la mémoire partagée. La boucle interne fait 1 FMA pour 2 lectures partagées — le ratio est de 1:2, très défavorable.


2.5 Étape 3 — le pavage en registres, 1D

Chaque thread calcule plusieurs éléments de \(\mathbf{C}\) dans une colonne :

#define BM 64
#define BN 64
#define BK  8
#define TM  8          // 8 éléments de C par thread, en colonne
// 64*64/8 = 512 threads par bloc

// Dans la boucle interne :
for (int k = 0; k < BK; ++k) {
    float b_val = Bs[k][col_thread];              // 1 lecture
    #pragma unroll
    for (int i = 0; i < TM; ++i)
        acc[i] += As[k][lig_thread*TM + i] * b_val;   // TM lectures, TM FMA
}

Ratio calcul/mémoire-partagée : \(TM\) FMA pour \(TM + 1\) lectures, soit ~1:1. Amélioration nette, insuffisante.

Gain : facteur ~2. On atteint ~50 % du pic.


2.6 Étape 4 — le pavage en registres, 2D

L'étape décisive. Chaque thread calcule un bloc \(TM \times TN\) de \(\mathbf{C}\) :

#define BM 128
#define BN 128
#define BK   8
#define TM   8
#define TN   8
// 128*128/(8*8) = 256 threads par bloc

float acc[TM][TN] = {0.0f};
float regA[TM], regB[TN];

for (int k = 0; k < BK; ++k) {
    #pragma unroll
    for (int i = 0; i < TM; ++i) regA[i] = As[k][ty*TM + i];
    #pragma unroll
    for (int j = 0; j < TN; ++j) regB[j] = Bs[k][tx*TN + j];

    #pragma unroll
    for (int i = 0; i < TM; ++i)
        #pragma unroll
        for (int j = 0; j < TN; ++j)
            acc[i][j] += regA[i] * regB[j];        // produit extérieur
}

Le produit extérieur : \(TM + TN = 16\) lectures en mémoire partagée pour \(TM \times TN = 64\) FMA. Ratio 4:1, contre 1:2 à l'étape 2.

Gain : facteur ~1,6. On atteint 75 à 85 % du pic.

Coût : ~100 registres par thread, donc une occupancy de ~25 %. C'est le régime Volkov, et c'est délibéré.

L'idée centrale de toute l'optimisation de GEMM

Il y a trois niveaux de pavage :

Niveau Support Ce qu'on y garde
Bloc mémoire partagée tuile \(BM \times BK\) et \(BK \times BN\)
Thread registres bloc \(TM \times TN\) d'accumulateurs
Instruction tensor core fragments MMA

Chaque niveau réduit le trafic vers le niveau supérieur. C'est exactement la hiérarchie mémoire de la partie 1, exploitée systématiquement.


2.7 Étape 5 — vectorisation et transposition de As

Deux améliorations qui vont ensemble.

Chargement vectorisé :

float4 tmp = reinterpret_cast<const float4*>(&A[...])[0];

Une instruction LDG.128 au lieu de quatre LDG.32.

Stockage transposé de As : on range As en [BK][BM] au lieu de [BM][BK], de sorte que la lecture d'une colonne dans la boucle interne devienne contiguë — et donc vectorisable en float4 elle aussi, sans conflit de banc.

Gain : facteur ~1,1. On approche 90 % du pic FP32 vectoriel.


2.8 Étape 6 — le double tampon

Charger la tuile \(t+1\) pendant qu'on calcule la tuile \(t\) (voir CUDA moderne 1).

Gain : facteur ~1,05 à 1,1. Autour de 95 % du pic FP32.

C'est le point où Simon Boehm s'arrête : « à moins de 5 % de cuBLAS » en FP32, à partir d'un noyau naïf, en une dizaine d'étapes documentées.


2.9 Étape 7 — les tensor cores

Toutes les étapes précédentes concernent le FP32 vectoriel. Passer aux tensor cores change d'échelle : le pic passe de 67 à 990 TFLOPS en BF16.

Le chemin :

Approche Fraction du pic BF16
API wmma ~55-60 %
mma PTX manuel avec ldmatrix ~63 %
wgmma + TMA + warp specialization ~95 %

Le chiffre de 62,9 % pour mma contre 95 % pour wgmma est mesuré dans le microbenchmarking de Hopper.

Écrire l'étape 7 à la main est un projet de plusieurs mois. C'est là qu'on utilise CUTLASS, CuTe DSL ou Triton.


2.10 Étape 8 — les détails qui restent

Ce qui sépare 90 % de 98 %, et que fait cuBLAS :

Technique Effet
Ordonnancement des tuiles (regroupement) localité L2, +10-30 %
Éviter la quantification de vagues choisir des tailles de tuile qui divisent la grille
Split-K / Stream-K pour les matrices « hautes et fines » où \(M \cdot N\) est petit
Multicast TMA dans un cluster divise le trafic HBM des opérandes partagés
Épilogue fusionné biais, activation, quantification sans passe supplémentaire
Sélection d'algorithme une centaine de noyaux, choisis selon \(M, N, K\)

Le dernier point est décisif : cuBLAS ne contient pas un noyau GEMM mais des centaines, spécialisés par forme, précision et architecture, avec une heuristique de sélection. C'est ce que vous ne reproduirez pas.

Stream-K mérite une mention : c'est une stratégie d'ordonnancement qui découpe le travail selon la dimension \(K\) pour équilibrer parfaitement la charge entre SM, éliminant la quantification de vagues. Elle est particulièrement efficace sur les formes irrégulières.


2.11 Le tableau récapitulatif

Étape Technique % du pic FP32 Facteur cumulé
0 Naïve 1,3 % 1×
1 Coalescence ~10 % 8×
2 Mémoire partagée ~30 % 23×
3 Registres 1D ~50 % 38×
4 Registres 2D ~80 % 62×
5 Vectorisation ~88 % 68×
6 Double tampon ~93 % 72×
7 Tensor cores (wgmma) 95 % du pic BF16 —
8 Ordonnancement, algo 98 % —

Ces pourcentages sont indicatifs

Ils correspondent aux ordres de grandeur rapportés dans les worklogs publics (Simon Boehm sur A6000, salykova, cudaforfun sur H100) et varient selon la carte, les dimensions et la version du compilateur. Ils n'ont pas été reproduits ici.

Ce qui est robuste : l'ordre des étapes et la taille relative des sauts. Les étapes 1, 2 et 4 apportent l'essentiel.


2.12 Ce qu'il faut en retenir

Les cinq leçons transférables

  1. La coalescence d'abord. Facteur 8, pour deux lignes de code.
  2. La mémoire partagée transforme le régime. \(I\) passe de \(1/4\) à \(T/4\).
  3. Les registres sont le vrai levier. Le produit extérieur donne \(TM \cdot TN\) FMA pour \(TM + TN\) lectures. C'est l'étape 4 qui fait la différence.
  4. Une occupancy basse est normale et souhaitable dans ce régime.
  5. Vous n'écrirez pas les étapes 7 et 8. Utilisez CUTLASS, Triton ou cuBLAS.

Et surtout : cette séquence est le modèle de toute optimisation de noyau dense. FlashAttention est structurellement la même chose appliquée à \(\mathbf{Q}\mathbf{K}^\top\) puis à \(\operatorname{softmax}\cdot\mathbf{V}\).


Vérifiez que vous avez compris

Pourquoi le passage de TM=8 à TM=16 n'améliore-t-il pas indéfiniment les performances ?

Parce que acc[16][16] = 256 registres, au-delà de la limite matérielle de 255 par thread. Le compilateur déverse en mémoire locale (spilling), qui est de la mémoire globale : le noyau s'effondre.

Même avant cette limite, augmenter \(TM\) et \(TN\) réduit le nombre de warps résidents, donc la capacité à cacher la latence des chargements de tuiles.

L'optimum se situe généralement à \(TM = TN = 8\) (64 accumulateurs) ou \(8 \times 4\). C'est le résultat d'un compromis, pas d'une monotonie.

Une GEMM avec M=4096, N=4096, K=16. Que se passe-t-il ?

C'est une matrice « plate » : beaucoup de sortie, très peu de profondeur.

  • Intensité arithmétique : \(2MNK / (2(MK + KN + MN)) \approx 2 \times 16 / 2 \approx 16\) — très en dessous du seuil, la GEMM est limitée par la mémoire ;
  • Travail par tuile : avec des tuiles \(128\times128\times8\), il n'y a que 2 itérations de la boucle K. Le prologue et l'épilogue dominent ;
  • Nombre de blocs : \(32 \times 32 = 1024\), ce qui remplit bien le GPU.

C'est un cas où Split-K aide : découper la dimension \(K\) n'a pas de sens ici (elle est déjà minuscule), mais réduire la taille de tuile pour augmenter le nombre de blocs et mieux paralléliser oui.

Plus généralement : les GEMM à petit \(K\) sont limitées par la mémoire et exigent d'autres paramètres que les GEMM carrées. C'est une des raisons de la centaine de noyaux de cuBLAS.

Vous écrivez une GEMM en Triton et obtenez 60 % de cuBLAS. Où chercher les 40 % ?

Dans l'ordre de rentabilité :

  1. L'autotuning — avez-vous exploré BLOC_M/N/K, num_warps et num_stages ? C'est souvent 20 % à lui seul.
  2. Le regroupement de blocs (GROUPE_M) — 10 à 30 % de localité L2.
  3. La forme — cuBLAS a un noyau spécialisé pour vos dimensions, Triton en a un générique. Sur des dimensions inhabituelles, l'écart se réduit souvent.
  4. tl.dot émet-il bien wgmma ? Vérifiez le PTX généré.
  5. La quantification de vagues — vos dimensions produisent-elles un nombre de blocs qui est un multiple des blocs par vague ?

Et la question de fond : 60 % de cuBLAS est-il suffisant ? Si votre GEMM n'est qu'une partie du pipeline et que Triton vous permet de la fusionner avec l'épilogue, le gain de fusion peut dépasser les 40 % perdus.


Chapitre suivant : 3 · FlashAttention


Sources de ce chapitre