Almost every FLOP in training and serving a transformer is spent in general matrix multiplication, GEMM: the QKV and output projections, the MLP layers, the LM head, and the matching gradient computations in the backward pass. A GPU can execute GEMM at close to its peak throughput, but only because the algorithm is organised very carefully around the memory hierarchy.

This article builds a GPU matrix multiply in four levels, each one fixing the bottleneck of the previous: one thread per output, shared-memory tiles, register tiles, and finally tensor cores. For each level it works out the arithmetic intensity, the number of FLOPs performed per byte moved, and compares it with what the hardware needs. The goal is not to beat cuBLAS but to understand why its kernels look the way they do and which shapes run fast.

Advertisement

The problem and the roofline

C = A x B with A of shape M x K and B of shape K x N performs M x N x K multiply-adds, or 2MNK floating-point operations. It touches at least (MK + KN + MN) elements. For square 4,096 x 4,096 matrices in BF16 that is about 1.37 x 1011 FLOPs over at least 100 MB, an arithmetic intensity of roughly 1,365 FLOPs per byte if each element crosses HBM exactly once. That is far above the ridge point of any GPU, so large GEMMs are compute-bound in principle.

On an H100 SXM, dense BF16 tensor-core throughput is about 989 TFLOP/s and HBM bandwidth about 3.35 TB/s, a ridge of roughly 295 FLOPs per byte. For FP32 on the CUDA cores, 67 TFLOP/s over the same bandwidth gives a ridge near 20. A kernel whose effective intensity against HBM falls below the ridge is memory-bound no matter how fast its arithmetic is. The whole game in GEMM is to make the intensity each level of the hierarchy sees as high as possible by reusing data held closer to the arithmetic units.

Level 0: one thread per output

The direct mapping gives each thread one element of C and lets it loop over K. It is correct, easy, and slow. Each iteration loads one element of A and one of B, 8 bytes in FP32, for one FMA, 2 FLOPs, so the intensity against global memory is 0.25 FLOPs per byte. Caches recover some reuse, but the kernel stays memory-bound.

// Level 0: one thread per output element. Row-major A (M x K), B (K x N), C (M x N).
__global__ void gemm_naive(const float* A, const float* B, float* C, int M, int N, int K) {
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    int col = blockIdx.x * blockDim.x + threadIdx.x;   // x fastest -> adjacent columns
    if (row < M && col < N) {
        float acc = 0.f;
        for (int k = 0; k < K; ++k)
            acc += A[row * K + k] * B[k * N + col];     // 2 loads per FMA
        C[row * N + col] = acc;
    }
}

The thread mapping already matters here. With col taken from threadIdx.x, the 32 lanes of a warp read 32 consecutive elements of a B row, which coalesces, and all read the same A element, which is a broadcast. SIMT execution explains why the warp is the unit that matters for these accesses.

Advertisement

Level 1: shared-memory tiling

The GEMM data path: each level reuses data loaded by the level aboveHBM / L2A (M x K), B (K x N), C (M x N)Block tile in shared memoryBM x BK of A, BK x BN of BThread / warp tile in registersTM x TN accumulatorsread once per block tilereused BN (or BM) timesreused TN (or TM) timesTensor core MMA16 x 16 x 16 per warp stepEpiloguescale, bias, activation, castC tile written onceafter all K stepsGrid of blocks over C: block (bx, by) owns rows by*BM .. +BM and columns bx*BN .. +BNThe highlighted tile loops over K in BK slices: load A and B slices, sync, multiply, sync, repeat.
Each thread block owns a tile of C and marches along K, staging slices of A and B in shared memory; threads accumulate their outputs in registers, and tensor cores consume warp-sized fragments.

The fix for redundant loads is to share them. A block of T x T threads computes a T x T tile of C. It walks along K in steps of T: all threads cooperatively load a T x T tile of A and one of B into shared memory, synchronize, and each thread performs T multiply-adds reading only from shared memory. Every element loaded from global memory is then used T times instead of once.

The arithmetic: per K step the block loads 2T2 floats, 8T2 bytes, and performs T3 FMAs, 2T3 FLOPs. Intensity is T/4 FLOPs per byte. With T = 32 that is 8, a 32x improvement over level 0 but still below the FP32 ridge of about 20 against HBM. The L2 cache helps more than this figure suggests, because many blocks read the same tiles of A and B, but shared-memory tiling alone does not reach peak.

// Level 1: T x T tiles staged in shared memory. Launch with blockDim = (T, T).
template <int T>
__global__ void gemm_tiled(const float* A, const float* B, float* C, int M, int N, int K) {
    __shared__ float As[T][T];
    __shared__ float Bs[T][T];
    int ty = threadIdx.y, tx = threadIdx.x;
    int row = blockIdx.y * T + ty, col = blockIdx.x * T + tx;
    float acc = 0.f;
    for (int k0 = 0; k0 < K; k0 += T) {
        // each thread loads one element of each tile (coalesced along x)
        As[ty][tx] = (row < M && k0 + tx < K) ? A[row * K + k0 + tx] : 0.f;
        Bs[ty][tx] = (k0 + ty < K && col < N) ? B[(k0 + ty) * N + col] : 0.f;
        __syncthreads();                               // tile fully loaded
        #pragma unroll
        for (int k = 0; k < T; ++k)
            acc += As[ty][k] * Bs[k][tx];              // T FMAs from shared memory
        __syncthreads();                               // before tile is overwritten
    }
    if (row < M && col < N) C[row * N + col] = acc;
}

Two barriers are required per step: one so that no thread reads a tile before it is fully written, and one so that no thread overwrites it while others still read. In this layout a warp reads one row of Bs and broadcasts one element of As, so there are no bank conflicts; kernels that read a tile column-wise pad it by one column, and GPU shared memory works through the bank arithmetic.

Level 2: register tiling and the block-tile shape

Level 1 has a hidden limit: every FMA still reads two operands from shared memory, and shared-memory bandwidth, while much higher than HBM, is not unlimited. The next step is to give each thread several outputs. A thread owning a TM x TN patch of C loads TM values of A and TN values of B from shared memory per k and performs TM x TN FMAs using registers only. With TM = TN = 8, that is 64 FMAs for 16 shared loads, an eightfold cut in shared loads per FMA compared with one output per thread.

At the same time the block tile grows and becomes rectangular. A block computing BM x BN with a K slice of BK loads (BM + BN) x BK elements per step and performs BM x BN x BK FMAs, so its intensity against global memory is 2 x BM x BN / ((BM + BN) x bytes per element). For BM = BN = 128 in BF16 that is 64 FLOPs per byte from the block's own loads, with L2 hits on tiles shared between blocks raising the effective figure further.

# Level 2 pseudocode: block tile BM x BN, K slice BK, each thread owns TM x TN outputs.
for each block (by, bx) in parallel:
    acc[TM][TN] = 0                                   # registers, per thread
    for k0 in range(0, K, BK):
        cooperatively load A[by*BM : +BM, k0 : +BK] -> As   (shared)
        cooperatively load B[k0 : +BK, bx*BN : +BN] -> Bs   (shared)
        barrier()
        for k in range(BK):
            a_frag[0:TM] = As[my_rows, k]             # TM shared loads
            b_frag[0:TN] = Bs[k, my_cols]             # TN shared loads
            for i in range(TM):
                for j in range(TN):
                    acc[i][j] += a_frag[i] * b_frag[j]   # TM*TN FMAs
        barrier()
    epilogue: C[my_rows, my_cols] = alpha * acc + beta * C   (+ bias, activation)

Production kernels add two more ideas at this level. Double buffering loads the next K slice into a second shared-memory buffer while computing on the current one, so memory latency overlaps with arithmetic; newer GPUs provide asynchronous copy hardware for this.

Level 3: tensor cores

On FP32 CUDA cores a register-tiled kernel can approach the FP32 peak, but training runs in BF16 or FP16 on tensor cores, which multiply small matrix fragments per instruction. The portable CUDA interface is WMMA, warp-level matrix multiply-accumulate in the nvcuda::wmma namespace, available from compute capability 7.0. A warp declares fragments, loads them with load_matrix_sync, multiplies with mma_sync and writes back with store_matrix_sync. For half-precision inputs the supported shapes are m16n16k16, m32n8k16 and m8n32k16, with a float or half accumulator.

#include <mma.h>
using namespace nvcuda;

// Level 3: one warp computes a 16x16 tile of C with tensor cores.
// A row-major (M x K) half, B column-major (K x N) half, C float. Dims multiples of 16.
__global__ void gemm_wmma(const half* A, const half* B, float* C, int M, int N, int K) {
    int warp_m = (blockIdx.y * blockDim.y + threadIdx.y);
    int warp_n = (blockIdx.x * blockDim.x + threadIdx.x) / 32;
    wmma::fragment<wmma::matrix_a, 16, 16, 16, half, wmma::row_major> a;
    wmma::fragment<wmma::matrix_b, 16, 16, 16, half, wmma::col_major> b;
    wmma::fragment<wmma::accumulator, 16, 16, 16, float> acc;
    wmma::fill_fragment(acc, 0.0f);
    for (int k = 0; k < K; k += 16) {
        wmma::load_matrix_sync(a, A + warp_m * 16 * K + k, K);   // ld = K
        wmma::load_matrix_sync(b, B + warp_n * 16 * K + k, K);   // col-major, ld = K
        wmma::mma_sync(acc, a, b, acc);                           // acc += a * b
    }
    wmma::store_matrix_sync(C + warp_m * 16 * N + warp_n * 16, acc, N, wmma::mem_row_major);
}
// Teaching kernel: it loads fragments straight from global memory. Real kernels
// stage block tiles through shared memory first, exactly as in levels 1 and 2.

The fragment is opaque: its mapping of matrix elements to lanes is unspecified, and you are not supposed to depend on it. Libraries such as CUTLASS and cuBLAS go further, using lower-level matrix instructions and architecture-specific asynchronous features to keep tensor cores fed, but the algorithm is the same hierarchy: block tiles in shared memory, warp tiles as tensor-core fragments, accumulators in registers. The high tensor-core peak makes data movement the bottleneck once more, which is why the ridge rises from about 20 to about 295 FLOPs per byte and why double buffering and large tiles are mandatory rather than optional. GPU tensor cores covers the hardware side.

Shapes that run slowly, and why

  • Dimensions that are not multiples of the tile. A GEMM with N = 4,100 needs a partial tile column that does mostly padding work. Keep hidden sizes, head dimensions and vocabulary sizes multiples of 64 or 128 where you can; padding the vocabulary is a common, cheap fix.
  • Wave quantization. The grid is executed in waves of blocks across the SMs. With 132 SMs and one block per SM, a GEMM producing 140 tiles runs one full wave and a second wave with 8 blocks, so the device sits mostly idle for half the time. Small changes in tile size or batch can move a layer across such a boundary.
  • Skinny GEMMs. When M is small, as in decode-time inference with batch 1, the intensity collapses toward the matrix-vector case and the kernel is memory-bound however well it is tiled. The remedy is batching, not a better kernel.
  • Small K with large M and N gives few K steps to amortize loading and writing C. Large K with small M and N gives too few blocks to fill the GPU; split-K fixes that by dividing K among blocks and summing partial results, at the cost of an extra reduction and non-deterministic summation order if atomics are used.
  • FP32 by default. Without BF16, FP16 or TF32 enabled, matmuls do not use tensor cores at all on many configurations; in PyTorch, torch.backends.cuda.matmul.allow_tf32 and autocast control this.

From one GPU to many

The same tiling idea scales beyond one device. Tensor parallelism splits A or B by columns or rows across GPUs so that each computes a block of C, followed by an all-gather or reduce-scatter, exactly as blocks on one GPU split C. Classical algorithms such as Cannon and SUMMA rotate or broadcast blocks of A and B between processors arranged in a grid. Computation per tile grows with the cube of the tile edge, communication with its square. Model parallel training shows how this appears in practice.

Operational guidance and trade-offs

Use cuBLAS, cuBLASLt or a framework's matmul in production; they select among many tuned kernels per shape. Your levers are the shapes and dtypes you give them. Profile a training step, list the top GEMMs by time, compute each one's FLOPs and divide by duration to get achieved throughput, and compare with the dtype's peak. GEMMs well below peak usually have a shape problem, a precision problem, or a missing fusion. Fuse cheap element-wise work, such as bias, activation and casts, into the GEMM epilogue while outputs are still in registers; see kernel fusion. Custom kernels pay off for fused patterns the libraries do not cover, such as attention or quantized GEMM with unusual layouts, and tools like Triton make the tiling above expressible in tens of lines.

The central trade-off inside the kernel is tile size. Bigger tiles mean more reuse and higher intensity but more registers and shared memory per block, fewer resident blocks, coarser wave quantization and worse handling of edge tiles. Smaller tiles fit better and fill the machine on small problems but move more data per FLOP.

Key takeaway: A fast GPU matmul is a data-reuse hierarchy. Shared-memory tiles cut global traffic by the tile edge, register tiles cut shared-memory traffic by the per-thread patch, and tensor cores raise the arithmetic peak so high that double buffering and large tiles become mandatory. Arithmetic intensity at each level tells you which one is the bottleneck. In practice, rely on library kernels, pick shapes that are multiples of the tile sizes, run in BF16 or FP16, fuse epilogues, and check achieved TFLOP/s against peak for your largest GEMMs.