gpu-kernels-explained
← /learn · 07

GEMM, step by step

Naive, tiled in shared memory, register-blocked, tensor cores: each step reuses data closer to the ALUs and raises the arithmetic intensity.

Loading the animation…

Concept

A matrix multiply C=ABC = AB does 2n32n^3 flops on 3n23n^2 numbers, so in principle every number can be reused nn times. A kernel's whole job is to arrange that reuse close to the ALUs. The animation shows four ways, on a problem small enough to draw: a 16 × 16 × 16 multiply, worked a block of C at a time, with the k dimension in slices of 4.

  1. Naive. One thread per output; each thread reads its row of A and its column of B straight from global memory. Every value of A is fetched once per column of C. The counters stop at 0.25 flop/byte on the 4096³ problem: the GPU spends its time waiting.
  2. Tiled in shared memory. A block of threads loads a tile of A and a tile of B into shared memory together, once, and every thread of the block reads them from there. A value fetched from global memory is now used by a whole row or column of the block: intensity at L2 rises to 7.97 flop/byte for 32 × 32 blocks. But each multiply-add still reads two values from shared memory, so shared memory becomes the limit (chapter 1).
  3. Register blocking. Each thread computes an 8 × 8 patch of C: per k it loads 8 values of A and 8 of B into registers and does 64 multiply-adds with them. Shared-memory traffic per flop falls 8-fold, the block grows to 128 × 128, and on the A100 the kernel finally reaches the FP32 roof: 19.5 TFLOP/s in the model.
  4. Tensor cores. A warp hands whole tiles to a matrix instruction (mma.sync on the A100, wgmma on the H100) that does a small matrix multiply in one go, with BF16 inputs and FP32 sums. Half the bytes per value and a roof 16 times higher (on the A100). Now the same 128 × 128 tiles give 62.1 flop/byte from HBM, well below the tensor ridge at 201 flop/byte: under the model's no-reuse assumption the kernel is memory-bound again, at 96.5 TFLOP/s.

That last line is the lesson of the chapter: each step moved the bottleneck, and the faster the ALUs, the more reuse a kernel needs. Production tensor-core kernels use larger tiles (128 × 256), keep neighbouring blocks working on the same rows and columns so their loads hit L2, and on Hopper share tiles between blocks of a cluster.

Maths

Take a block that computes a bm×bnb_m \times b_n tile of C, stepping through k in slices of bkb_k. Each step loads bmbkb_m b_k values of A and bkbnb_k b_n of B (ee bytes each) and does 2bmbnbk2 b_m b_n b_k flops, so

Iglobal=2 bmbnbke (bmbk+bkbn)=2e (1/bn+1/bm),I_{\text{global}} = \frac{2\, b_m b_n b_k}{e\,(b_m b_k + b_k b_n)} = \frac{2}{e\,(1/b_n + 1/b_m)},

which for square blocks is b/eb/e: b/4b/4 for FP32 and b/2b/2 for BF16. The slice depth bkb_k cancels: it sets how much shared memory a stage needs, not the reuse.

Inside the block, a thread that owns a tm×tnt_m \times t_n patch reads tm+tnt_m + t_n values from shared memory per k and does 2tmtn2 t_m t_n flops with them, so its shared-memory intensity is

Ismem=2 tmtn4 (tm+tn),I_{\text{smem}} = \frac{2\, t_m t_n}{4\,(t_m + t_n)} ,

1/41/4 for one output per thread and 22 for an 8 × 8 patch (before counting the tile stores). That 8-fold gain is register blocking.

The registers pay for it: an 8 × 8 patch is 64 accumulators plus 16 operands per thread, so these kernels need well over 80 registers per thread, and occupancy falls. Fewer warps are needed, though, because each one has 64 independent multiply-adds per k to issue.

Code

The model's bytes for one k-tile of one block, cut from src/lib/gpu/model.ts:

if (staged) {
  const tile = eb * (bm * bk + bk * bn);
  glob += tile;
  smem += tile + eb * bk * idiv(bm * bn, tm * tn) * (tm + tn);
} else glob += eb * bm * bn * 2 * bk;

The shared-memory kernel's inner loop in CUDA, illustrative and not compiled by this site's CI:

for (int k0 = 0; k0 < K; k0 += TILE) {
  As[ty][tx] = A[row * K + k0 + tx];          // each thread loads one value
  Bs[ty][tx] = B[(k0 + ty) * N + col];
  __syncthreads();
  for (int k = 0; k < TILE; ++k)
    acc += As[ty][k] * Bs[k][tx];             // reused TILE times from smem
  __syncthreads();
}

The four kernels on a 4096³ GEMM (model)

KernelFlop per byte from L2Bound by (A100)Time, A100Time, H100
Naive0.25HBM354 ms164 ms
Shared-memory tiles (32²)7.97shared memory29.1 ms16.9 ms
Register blocking (128², 8²)31.5the FP32 ALUs7.05 ms2.18 ms
Tensor cores (BF16, 128²)62.1HBM (no L2 reuse)1.42 ms661 µs