1. What is it?#
- GEMM — GEneral Matrix Multiply:
C = α·A·B + β·C. Matrix times matrix. - GEMV — GEneral Matrix-Vector:
y = α·A·x + β·y. Matrix times vector.
These names come from BLAS (Basic Linear Algebra Subprograms), the 1979 interface that still governs how numerical libraries are organized. GEMM is “level 3” (O(n³) work on O(n²) data); GEMV is “level 2” (O(n²) work on O(n²) data).
That distinction is the whole point: GEMM has reuse, GEMV does not.
2. Why does it exist?#
Because prefill is GEMM and decode is GEMV, and they behave completely differently.
Prefill: (2048 tokens, 4096) @ (4096, 4096) → GEMM, arithmetic intensity ~1365
Decode: (1 token, 4096) @ (4096, 4096) → GEMV, arithmetic intensity ~1Same weights. Same math. A 1000x difference in efficiency. Understanding why is the single most important piece of mechanical knowledge in LLM inference.
3. Simple analogy#
A photocopier.
Copying one page: you walk to the machine, warm it up, place the page, copy, walk back. Ninety seconds for one page.
Copying 500 pages: same walk, same warm-up, then 500 copies. Ninety seconds plus copy time.
The walk and warm-up (loading the weight matrix from HBM) is fixed. GEMV pays it for one page. GEMM amortizes it over hundreds.
4. Tiny example#
// gemm.go — one weight matrix, M input rows: predict the GEMV → GEMM transition.
package main
import "fmt"
func main() {
const (
d = 4096.0
bytesPer = 2.0 // FP16
bandwidth = 950e9 // bytes/s (A100-class, achievable)
peak = 62e12 // FLOP/s (A100-class FP16 tensor cores, achievable)
)
for _, M := range []float64{1, 2, 8, 32, 128, 512, 2048} {
flops := 2 * M * d * d
bytes := bytesPer * (M*d + d*d + M*d) // read A, read W, write C
t := max(bytes/bandwidth, flops/peak)
fmt.Printf("M=%5.0f %8.1f us %7.2f TF/s %7.0f GB/s I=%7.1f\n",
M, t*1e6, flops/t/1e12, bytes/t/1e9, flops/bytes)
}
}Output (measured A100 numbers are within about 10% of this prediction):
M= 1 35.3 us 0.95 TF/s 950 GB/s I= 1.0
M= 2 35.4 us 1.90 TF/s 950 GB/s I= 2.0
M= 8 35.5 us 7.57 TF/s 950 GB/s I= 8.0
M= 32 35.9 us 29.93 TF/s 950 GB/s I= 31.5
M= 128 69.3 us 62.00 TF/s 515 GB/s I= 120.5
M= 512 277.1 us 62.00 TF/s 151 GB/s I= 409.6
M= 2048 1108.4 us 62.00 TF/s 61 GB/s I= 1024.0M=1 to M=8: identical time. You did 8x the work for free. That is the GEMV → GEMM transition, and it’s why batching is not optional.
5. Technical explanation#
Why GEMV cannot be fast#
GEMV: y = W x, W is (N, K), x is (K,)
FLOPs = 2·N·K
Bytes = 2·N·K (weights) + 2·K (input) + 2·N (output) ≈ 2·N·K
Intensity = 2NK / 2NK = 1 FLOP/byteEvery weight is read once and used once. There is no reuse to exploit — not a limitation of the implementation, a property of the operation. On an H100 with a ridge point of 296, you run at 1/296 of peak compute. No kernel can fix this.
The only escapes:
- Batch it (make M > 1) → becomes GEMM.
- Read fewer bytes → quantization.
- Use the weights for more tokens per read → speculative decoding, MoE (fewer weights per token).
Those three are Sections V.09/VII.02/VII.13. Everything follows from this one paragraph.
How a GEMM kernel is structured#
LEVEL 1 — Thread block tile (e.g. 128×128 output)
Each block computes a 128×128 patch of C
Loops over K in chunks of 32 or 64
Stages A and B tiles into shared memory
LEVEL 2 — Warp tile (e.g. 64×64)
Each of 4-8 warps handles a sub-tile
Loads fragments from shared memory into registers
LEVEL 3 — Instruction tile (e.g. m16n8k16)
Tensor core MMA instruction
Accumulates into registers in FP32
EPILOGUE
Apply bias, activation, scaling, quantization — all in registers
Write C to HBM onceKey techniques:
- Double buffering: load tile k+1 while computing tile k.
- Async copy (
cp.asyncon Ampere, TMA on Hopper): DMA from HBM to shared memory without occupying registers. - Swizzled shared memory layout: avoid bank conflicts.
- Split-K: for small M/N and large K, split the K loop across blocks and reduce — important for decode-shaped GEMMs.
Tile and wave quantization#
TILE QUANTIZATION
M=129 with a 128-tile → 2 tile-rows, second is 1/128 utilized
Effective utilization: 129/256 = 50%
WAVE QUANTIZATION
Grid of 133 blocks on 132 SMs → 2 waves, second has 1 block
Effective utilization: 133/264 = 50%Both produce the sawtooth pattern you see when plotting GEMM time vs M. Round your dimensions.
Batched and grouped GEMM#
Batched GEMM: many independent same-shape GEMMs (e.g. per-head attention)
cublasGemmStridedBatched
Grouped GEMM: many DIFFERENT-shape GEMMs in one launch
essential for MoE (each expert gets a different number of tokens)
and for multi-LoRA servingGrouped GEMM is a relatively recent addition to CUTLASS and is what makes efficient MoE serving possible (Section XIII.03).
6. Under the hood#
Which kernel gets selected matters enormously:
export CUBLASLT_LOG_LEVEL=5 # verbose: shows heuristic choices
# or use Nsight Compute to see the kernel name:
# ampere_h16816gemm_128x128_ldg8_stages_64x3_tn
# ^arch ^mma shape ^tile ^stagesReading that name tells you the tile size and pipeline depth chosen. If you see a generic
gemm_kernel or a _nn_ variant where you expected _tn_, a transpose is being handled
inefficiently.
For decode, cuBLAS often falls back to a GEMV kernel or a split-K GEMM. Some engines ship hand-written decode GEMM kernels (e.g. for W4A16) because the general library isn’t tuned for M=1..64 with quantized weights.
7. Performance implications#
| Shape | Regime | Achieved fraction of peak |
|---|---|---|
| M=1 (decode, batch 1) | memory bound | ~0.3% of FLOPs, ~90% of bandwidth |
| M=8-64 (typical decode) | memory bound | 2-20% of FLOPs |
| M=128-512 | transitional | 40-70% |
| M≥1024 (prefill) | compute bound | 70-90% |
Practical rule: if your decode achieves >90% of theoretical bandwidth, your GEMMs are as good as they can be — go optimize something else (or reduce bytes).
8. Production implications#
- Fuse QKV into one GEMM and gate+up into one GEMM. Fewer, larger GEMMs are more efficient and reduce launch count.
- Pad dimensions to tile-friendly multiples.
- Use engine-provided quantized GEMM kernels (Marlin, Machete, exllama kernels) rather than generic dequantize-then-GEMM — they are 2-3x faster for W4A16 decode.
- Check kernel selection after any upgrade. Library heuristics change and can regress.
- Split-K matters for decode. If your engine doesn’t use it, small-M GEMMs leave SMs idle.
9. Common mistakes#
Trying to optimize GEMV’s FLOPs. It’s memory bound. Reduce bytes instead.
Not fusing QKV/gate-up. Three small GEMMs instead of one large one: more launches, worse tiles.
Ignoring tile quantization. M=257 vs M=256 can cost 50%.
Writing your own GEMM for production. You will not beat CUTLASS.
Assuming cuBLAS picks the best kernel. Its heuristics are good but not perfect, especially
for unusual shapes. cublasLtMatmulAlgoGetHeuristic with autotuning exists for a reason.
10. Hands-on exercise#
A. Reproduce the table. Run the benchmark in section 4 on your GPU. Plot TFLOP/s and GB/s vs
M on log axes. Mark the ridge point. Record in numbers.md.
B. Tile quantization. Sweep M from 120 to 140 in steps of 1 and plot time. Find the jumps. Repeat for N.
C. Fusion win. Time three separate GEMMs (M,4096)@(4096,4096) ×3 vs one
(M,4096)@(4096,12288). Report the difference at M=1 and M=2048. Explain both.
D. Inspect kernel selection. Use Nsight Compute to find which cuBLAS kernel runs for M=1, M=64, M=2048. Decode the kernel names.
E. Split-K. For a shape with small M and large K, compare cuBLAS default vs an explicitly split-K algorithm. Measure the difference.
11. Interview questions#
- What is the arithmetic intensity of GEMV, and why can’t it be improved?
- Describe the tiling hierarchy of a modern GEMM kernel.
- What is tile quantization? Wave quantization? How do you avoid them?
- Why does decode use GEMV-shaped operations, and what are the three ways out?
- What is split-K and when does it help?
- What is grouped GEMM and which workloads need it?
- Your decode achieves 92% of theoretical HBM bandwidth. Is there anything left to optimize?
12. Further reading#
- [ESTABLISHED] Goto & van de Geijn, “Anatomy of High-Performance Matrix Multiplication”
- [REFERENCE] NVIDIA CUTLASS repository and its GEMM API documentation
- [REFERENCE] NVIDIA “Matrix Multiplication Background User’s Guide”
- [ESTABLISHED] Marlin kernel (IST-DASLab) — state-of-the-art W4A16 GEMM
- Next: 04 — Convolutions