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 does flops on numbers, so in principle every number can be reused 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.
- 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.
- 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).
- 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.
- Tensor cores. A warp hands whole tiles to a matrix instruction (
mma.syncon the A100,wgmmaon 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 tile of C, stepping through k in slices of . Each step loads values of A and of B ( bytes each) and does flops, so
which for square blocks is : for FP32 and for BF16. The slice depth cancels: it sets how much shared memory a stage needs, not the reuse.
Inside the block, a thread that owns a patch reads values from shared memory per k and does flops with them, so its shared-memory intensity is
for one output per thread and 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)
| Kernel | Flop per byte from L2 | Bound by (A100) | Time, A100 | Time, H100 |
|---|---|---|---|---|
| Naive | 0.25 | HBM | 354 ms | 164 ms |
| Shared-memory tiles (32²) | 7.97 | shared memory | 29.1 ms | 16.9 ms |
| Register blocking (128², 8²) | 31.5 | the FP32 ALUs | 7.05 ms | 2.18 ms |
| Tensor cores (BF16, 128²) | 62.1 | HBM (no L2 reuse) | 1.42 ms | 661 µs |