12. Fast matrix multiplication
In this chapter
- Data reuse: why the naive kernel is slow, and how tiling into shared memory fixes it.
- Barriers, boundary tiles and shared-memory bank conflicts.
- Register tiling (thread coarsening) and tensor cores: how libraries reach most of peak.
- Why decode's matrix-vector products are a different problem from prefill's GEMMs.
You will build
engine/tiling.py: a tiled matmul that counts its memory traffic (runs anywhere), and the tiled CUDA kernel in engine/kernels/cuda_ops.cu (needs a GPU).
Time: 6-8 hours. GPU: optional.
Why matrix multiplication deserves a chapter
In a transformer, nearly all arithmetic is matrix multiplication: the QKV, output and MLP projections of every layer, and the head. Prefill is a sequence of large GEMMs (general matrix-matrix multiplications), and their efficiency sets time-to-first-token. You’ll rarely beat NVIDIA’s cuBLAS or CUTLASS for large GEMMs, and you shouldn’t try. But every kernel you’ll write later (FlashAttention, quantized matmuls, paged attention) is built from the same tiling ideas. This chapter is where you learn them.
Counting the waste
Chapter 11’s naive kernel computes each output with $2K$ loads from global memory. For $M = N = K = 4096$, that’s $2 \times 4096^3 \approx 137$ billion loads to compute 16.8 million outputs, while the matrices themselves contain only $3 \times 4096^2 \approx 50$ million distinct values. Every element of A and B is fetched 4,096 times. The L2 cache catches some of the repeats, but the arithmetic units still mostly wait.
Tiling: load once, use many times
Divide the output into $T \times T$ tiles and give each tile to one thread block. To compute its tile, the block needs a $T$-row strip of A and a $T$-column strip of B. Walk along $K$ in phases. In each phase, the block cooperatively copies one $T\times T$ tile of A and one of B into shared memory, one element per thread, and then every thread accumulates $T$ products from those tiles:
Each loaded value is now used $T$ times instead of once, so global traffic drops by a factor of $T$:
$$ \text{loads}{\text{tiled}} = \frac{2MNK}{T}, \qquad I{\text{phase}} = \frac{2T^3 \text{ FLOPs}}{2T^2 \times 4 \text{ bytes}} = \frac{T}{4} \text{ FLOP/byte (FP32)}. $$
$T = 16$ gives an intensity of 4 instead of 0.25. That’s a 16x higher ceiling on the roofline.
You can verify the traffic arithmetic without a GPU. tiled_matmul performs the same tile-by-tile computation on the CPU and counts every element it “loads into shared memory”:
def tiled_matmul(a, b, tile=16):
"""C = A @ B computed one TILE x TILE output tile at a time, the way a CUDA block does it.
Returns (C, elements_loaded): every tile of A and B copied into "shared memory" is counted
as a global-memory load, so you can compare traffic with the naive kernel's 2*M*N*K. (Your engine: Chapter 12)
"""
m, k = a.shape
n = b.shape[1]
c = torch.zeros(m, n, dtype=torch.float32)
loads = 0
for row in range(0, m, tile): # one "block" per output tile
for col in range(0, n, tile):
acc = torch.zeros(min(tile, m - row), min(tile, n - col))
for phase in range(0, k, tile): # march along K one tile at a time
a_tile = a[row:row + tile, phase:phase + tile].float() # load into shared memory
b_tile = b[phase:phase + tile, col:col + tile].float()
loads += a_tile.numel() + b_tile.numel()
acc += a_tile @ b_tile # reuse each value `tile` times
c[row:row + tile, col:col + tile] = acc
return c, loads
For 64×64 matrices, it loads exactly $1/T$ as many elements as the naive kernel, for every tile size. That’s what the milestone tests check.
The CUDA kernel
constexpr int TILE = 16;
// Each block computes a TILE x TILE output tile. In phase p every thread loads one element of
// A's tile and one of B's into shared memory; then all threads reuse those 2*TILE^2 values
// TILE times each. Global traffic drops by a factor of TILE.
__global__ void tiled_matmul_kernel(const float* a, const float* b, float* c, int M, int K, int N) {
__shared__ float As[TILE][TILE];
__shared__ float Bs[TILE][TILE];
int row = blockIdx.y * TILE + threadIdx.y;
int col = blockIdx.x * TILE + threadIdx.x;
float total = 0.f;
for (int phase = 0; phase < K; phase += TILE) {
int ak = phase + threadIdx.x, bk = phase + threadIdx.y;
As[threadIdx.y][threadIdx.x] = (row < M && ak < K) ? a[row * K + ak] : 0.f; // zero-pad edges
Bs[threadIdx.y][threadIdx.x] = (bk < K && col < N) ? b[bk * N + col] : 0.f;
__syncthreads(); // tile fully loaded before anyone reads it
for (int k = 0; k < TILE; ++k) total += As[threadIdx.y][k] * Bs[k][threadIdx.x];
__syncthreads(); // everyone done reading before the next overwrite
}
if (row < M && col < N) c[row * N + col] = total;
}
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr):
"""(Your engine: Chapter 14)"""
# Map the 1-D program id to an output tile. Walking GROUP_M tile-rows at a time keeps the
# B tiles those rows need hot in L2 cache ("grouped ordering", Triton tutorial 03).
pid = tl.program_id(0)
tiles_m, tiles_n = tl.cdiv(M, BLOCK_M), tl.cdiv(N, BLOCK_N)
group = pid // (GROUP_M * tiles_n)
first_m = group * GROUP_M
group_size = min(tiles_m - first_m, GROUP_M)
pid_m = first_m + (pid % group_size)
pid_n = (pid % (GROUP_M * tiles_n)) // group_size
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
rk = tl.arange(0, BLOCK_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) # accumulate in FP32 whatever the inputs
for k0 in range(0, K, BLOCK_K):
a = tl.load(a_ptr + rm[:, None] * stride_am + (k0 + rk)[None, :] * stride_ak,
mask=(rm[:, None] < M) & ((k0 + rk)[None, :] < K), other=0.0)
b = tl.load(b_ptr + (k0 + rk)[:, None] * stride_bk + rn[None, :] * stride_bn,
mask=((k0 + rk)[:, None] < K) & (rn[None, :] < N), other=0.0)
acc += tl.dot(a, b) # tensor cores when dtypes allow
c = acc.to(c_ptr.dtype.element_ty)
tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn, c,
mask=(rm[:, None] < M) & (rn[None, :] < N))
def matmul(a, b, block_m=64, block_n=64, block_k=32, group_m=8):
check_device(a, b)
if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]:
raise ValueError("matmul expects [M, K] @ [K, N]")
m, k = a.shape
n = b.shape[1]
c = torch.empty((m, n), device=a.device, dtype=a.dtype)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n),)
matmul_kernel[grid](a, b, c, m, n, k, a.stride(0), a.stride(1), b.stride(0), b.stride(1),
c.stride(0), c.stride(1), BLOCK_M=block_m, BLOCK_N=block_n,
BLOCK_K=block_k, GROUP_M=group_m)
return c
Three details make the CUDA version correct:
- Two barriers per phase. The first
__syncthreads()ensures the whole tile is loaded before anyone reads it. The second ensures everyone has finished reading before the next phase overwrites the tile. Remove the second and fast threads corrupt the tile under slow ones. The result is wrong only sometimes, which makes it a nightmare to debug. - Zero-padded edges. When $M$, $N$ or $K$ isn’t a multiple of $T$, threads that would load out of bounds load 0 instead. Zeros don’t change the dot product, so edge tiles need no special arithmetic.
- Every thread reaches every barrier, including threads whose output is outside the matrix. They still help load tiles; they just skip the final write. A thread that returns early would make the barrier wait forever, or worse.
The Triton version (Chapter 14 explains the language) expresses the same algorithm at the level of tiles: tl.load a [BLOCK_M, BLOCK_K] block, tl.dot it with a [BLOCK_K, BLOCK_N] block, and accumulate in FP32. Triton handles shared memory and synchronization itself. It also maps tl.dot onto tensor cores, which our CUDA kernel doesn’t use.
Shared-memory bank conflicts
Shared memory is divided into 32 banks, and consecutive 4-byte words go to consecutive banks. A warp’s 32 accesses complete in one cycle if they hit 32 different banks (or the same word, which is broadcast). If several threads hit different words in the same bank, the accesses are serialized.
Reading a row of a float tile[32][32] across a warp touches 32 consecutive words: 32 banks, no conflict. Reading a column touches words 32 apart, all in the same bank: a 32-way conflict. The classic fix is padding: declare float tile[32][33]. Now advancing one row moves by 33 words, which is one bank over, so a column spans all 32 banks. Our tiled matmul reads As[ty][k] (a broadcast within a warp) and Bs[k][tx] (a row), so it’s conflict-free without padding. The transpose exercise in Chapter 11 needs the padding.
Register tiling and thread coarsening
Tiling into shared memory raised the intensity of global memory traffic. But each FMA in the inner loop still reads two values from shared memory, which has its own bandwidth limit. The next level of reuse moves into registers: each thread computes a small block of outputs, say $4\times4$ or $8\times8$, instead of one. Each value it reads from shared memory then feeds 4 or 8 FMAs held in registers:
for k in tile:
a_frag[0..7] = As[ty*8 .. ty*8+7][k] # 8 values of A
b_frag[0..7] = Bs[k][tx*8 .. tx*8+7] # 8 values of B
acc[i][j] += a_frag[i] * b_frag[j] # 64 FMAs from 16 loads
This is thread coarsening (PMPP §6.5 and Chapter 15): fewer threads, each doing more work, with higher reuse. The costs are more registers per thread (lower occupancy) and more complex code. Well-tuned FP32 SIMT kernels reach roughly 50-70% of FP32 peak this way.
Tensor cores
Modern NVIDIA GPUs have tensor cores: units that compute a small matrix product, such as a 16×8×16 BF16 tile with FP32 accumulation, as a single warp-wide instruction. Their throughput is roughly 8-16x that of the ordinary FP32 cores, which is where the “~990 TFLOP/s” of an H100 comes from. Using them requires BF16, FP16, FP8 or INT8 inputs, data arranged in specific fragment layouts, and enough reuse to feed them, so tiling matters even more.
You can program tensor cores with CUDA’s wmma/mma instructions (GPU Mode L23), with CUTLASS/CuTe templates (L15, L36, L57), or let Triton’s tl.dot do it. For the engine, the right choice is clear: use cuBLAS (torch.matmul) for the big GEMMs, and write kernels only where you can fuse something it can’t. Typical results for a 4096³ BF16 GEMM look like this:
| kernel | fraction of tensor-core peak |
|---|---|
| naive (Chapter 11), FP32 cores | under 1% |
| shared-memory tiled, FP32 cores | 2-5% |
| register-tiled, FP32 cores | 5-10% |
Triton tl.dot, BF16, autotuned | 60-85% |
| cuBLAS, BF16 | 70-90% |
(The fractions are relative to the tensor-core peak, which is why even a well-tuned FP32 SIMT kernel looks small here. Measure your own device in the stretch exercises.)
Decode is a different problem: GEMV
During decode, each projection multiplies a single vector (or a few, with batching) by a weight matrix. That’s a matrix-vector product (GEMV), with an arithmetic intensity of about 1 (Chapter 10). Tiling can’t help, because there’s nothing to reuse: each weight is used exactly once per token. A good GEMV kernel just streams the matrix at full bandwidth:
- coalesced, vectorized loads (each thread reads 16 bytes at a time, consecutive threads read consecutive memory);
- enough parallelism: one or more warps per output row, combining partial sums with warp reductions (Chapter 13);
- no wasted bytes: weights stored compactly, ideally in 4-bit form and dequantized in registers (Chapter 20).
That’s why decode optimization in practice is quantization plus bandwidth-efficient GEMV, and prefill optimization is GEMM. Your Rust engine’s matvec (rows split across CPU threads, BF16 widened in the inner loop) is a CPU GEMV in exactly this spirit.
Build it
Engine milestone 12: tiling. Implement tiled_matmul in engine/tiling.py. On a GPU, also write tiled_matmul_kernel in engine/kernels/cuda_ops.cu.
pytest tests/test_ch12_matmul.py # runs anywhere
pytest tests/test_ch11_cuda.py -k matmuls # GPU: naive and tiled vs torch
Then measure. Time naive_matmul, tiled_matmul and torch.matmul for 1024³ and 4096³ FP32 matrices (set torch.backends.cuda.matmul.allow_tf32 = False for a fair FP32 comparison) and convert the times to TFLOP/s.
Stretch exercises
- ★ Run
tiled_matmulwith tiles 4, 8, 16 and 32 on 256×256 matrices and plot loads against tile size. Then estimate the shared-memory capacity a 64×64 FP32 tile pair would need. Would it fit on your GPU? Where:experiments/ch12.py(create it), callingengine.tiling.tiled_matmulwith each tile size. - ★★ Add 2×2 register tiling to the CUDA kernel (each thread computes four outputs) and measure the speedup. Where:
tiled_matmul_kerneland its launch geometry inengine/kernels/cuda_ops.cu. - ★★ Write a GEMV kernel (
y = W x, W[N, K]BF16) where each warp computes one output row with a warp-shuffle reduction. Measure its achieved bandwidth against the datasheet. Where: add a GEMV kernel, host launcher and binding inengine/kernels/cuda_ops.cu. - ★★★ Use
wmmato write a tensor-core BF16 GEMM for multiples of 16, and compare it with cuBLAS. Where: add a WMMA kernel, host launcher and binding inengine/kernels/cuda_ops.cu.
Check your understanding
- Why does a tiled matmul need two barriers per phase, and what goes wrong without the second?
- Why can larger tiles be slower, even though they increase reuse?
- A warp reads column 5 of
float s[32][32]. How many bank conflicts occur, and how doess[32][33]fix it? - Why doesn’t tiling help a decode-time matrix-vector product?
Going deeper
- PMPP Chapter 5 (pp. 103-130): memory types, tiling, the tiled matmul kernel, boundary checks. §6.4-6.5: bank conflicts and coarsening. Chapter 15 (GEMM): register tiling, software pipelining and tensor-core considerations.
- GPU Mode L5 (Jeremy Howard, tiled matmul in CUDA and Numba from Python), L23 (tensor cores), L15/L36/L57 (CUTLASS and CuTe).
- Simon Boehm, How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance (blog, 2022): ten kernels from naive to 90% of cuBLAS, each measured.