PidokuInfra

Matrix Multiplication by Hand

Foundations Beginner 1h 15m Difficulty 2/5 Topic 02 of 12

Prerequisites 01

Matrix multiplication is 90-95% of the arithmetic in LLM inference. Everything in Sections VI, VII, and IX is, ultimately, about making this one operation faster. Do the exercises by hand.


1. What is it?#

Given A of shape (M, K) and B of shape (K, N), the product C = A @ B has shape (M, N), and

C[i][j] = sum over k of  A[i][k] * B[k][j]

In words: element (i,j) of the result is the dot product of row i of A with column j of B.

The inner dimensions must match. That’s the whole rule.

(M, K) @ (K, N) = (M, N)
      ↑     ↑
      must match; they vanish

2. Why does it exist?#

Because “apply a linear transformation to many vectors at once” is the fundamental operation of neural networks, and matrix multiplication is exactly that.

One vector through a layer:   y = W x            (M,K)@(K,1)
A batch of vectors:           Y = X @ Wᵀ         (B,K)@(K,M)

Composing layers is composing matrices. And because matmul has high arithmetic intensity (Section I.08), it is the one operation hardware designers have optimized relentlessly — tensor cores exist for it and nothing else.


3. Simple analogy#

A pricing table.

Rows of A = customers, columns of A = quantities of each product. Rows of B = products, columns of B = prices in different currencies.

C[i][j] = what customer i pays in currency j. Each entry combines “how much of each product” with “what each product costs” — a dot product.


4. Tiny example#

Do this on paper before reading the answer.

A = [1  2]      B = [5  6]
    [3  4]          [7  8]

C = A @ B

C[0][0] = row0(A) · col0(B) = 1·5 + 2·7 = 5 + 14 = 19
C[0][1] = row0(A) · col1(B) = 1·6 + 2·8 = 6 + 16 = 22
C[1][0] = row1(A) · col0(B) = 3·5 + 4·7 = 15 + 28 = 43
C[1][1] = row1(A) · col1(B) = 3·6 + 4·8 = 18 + 32 = 50

C = [19  22]
    [43  50]

Now the non-square case from the task description:

[1 2] × [3]   →  (1,2) @ (2,1) = (1,1)
        [4]

result = 1·3 + 2·4 = 11

And a shape you will meet constantly — a batch through a layer:

X = [1  0  2]      W = [ 1  0]      X is (2,3), W is (3,2)
    [0  3  1]          [ 2  1]
                       [-1  1]

C[0][0] = 1·1 + 0·2 + 2·(-1) = -1
C[0][1] = 1·0 + 0·1 + 2·1    =  2
C[1][0] = 0·1 + 3·2 + 1·(-1) =  5
C[1][1] = 0·0 + 3·1 + 1·1    =  4

C = [-1  2]     shape (2,2)
    [ 5  4]

Two inputs (rows of X), each transformed from 3-dim to 2-dim. That is one layer of a neural network on a batch of 2.


5. Technical explanation#

Cost#

FLOPs = 2 · M · N · K            (one multiply + one add per term)

For a transformer’s FFN up-projection with d=4096, d_ff=14336, batch of 2048 tokens:

2 · 2048 · 14336 · 4096 = 2.4e11 = 240 GFLOPs      for ONE matrix in ONE layer

Multiply by 3 (SwiGLU has 3 matrices), by 32 layers: ~23 TFLOPs for the FFN alone in one prefill. On an H100 at ~600 TFLOP/s achieved, ~38 ms. That is where your TTFT goes.

Three ways to think about it (all useful)#

1. Dot products (the definition). Good for understanding.

2. Linear combinations of columns. C[:, j] = sum_k B[k][j] · A[:, k]. Each output column is a weighted mix of A’s columns. Good for intuition about what a layer does.

3. Sum of outer products. C = sum_k A[:, k] ⊗ B[k, :]. Each k contributes a rank-1 update. This is how hardware actually does it — tensor cores accumulate rank-1 (or small rank) updates into an accumulator held in registers.

Properties#

Associative:      (AB)C = A(BC)        ← but the FLOP counts differ enormously!
Distributive:     A(B+C) = AB + AC
NOT commutative:  AB ≠ BA
Transpose:        (AB)ᵀ = Bᵀ Aᵀ
Identity:         AI = IA = A

Associativity is a performance tool. Consider (A @ B) @ C with A:(1,4096), B:(4096,4096), C:(4096,4096):

(A@B)@C:  2·1·4096·4096 + 2·1·4096·4096 = 67 MFLOPs
A@(B@C):  2·4096·4096·4096 + 2·1·4096·4096 = 137 GFLOPs      2000x worse!

Same answer, 2000x the work. Compilers and libraries exploit this; so should you when writing custom code. (LoRA’s efficiency is exactly this trick: x @ (A @ B) where A and B are thin.)

Blocked/tiled multiplication#

Real implementations don’t compute one element at a time. They partition into tiles:

C[I,J] = sum over Kk of  A[I,Kk] @ B[Kk,J]

where I, J, Kk index TILES (e.g. 128×128 or 64×64)

Each tile of A and B is loaded into fast memory once and used for a whole tile’s worth of output. This raises arithmetic intensity from ~1 to ~tile_size/3 (Section II.02). Every production GEMM — cuBLAS, CUTLASS, oneDNN — is a carefully tuned tiling hierarchy:

Grid level:     which output tile does this thread block compute?
Block level:    tiles staged in shared memory
Warp level:     sub-tiles held in registers
Instruction:    tensor core MMA on 16×8×16 fragments

Section VI.08 has you write a simple version of this.


6. Under the hood#

A modern GEMM kernel on an H100, roughly:

for each output tile (128×128 assigned to a thread block):
    accumulator[128][128] in registers (distributed across warps)
    for k_tile in range(0, K, 64):
        async-copy A[128×64] and B[64×128] from HBM to shared memory   (TMA on Hopper)
        __syncthreads()
        for each warp:
            load fragments from shared memory to registers
            issue tensor-core MMA instructions accumulating into registers
    write accumulator to C in HBM (epilogue: may fuse bias, activation, quantization)

Key ideas visible here that recur throughout the curriculum: double buffering (load the next tile while computing the current), accumulate in registers (avoid HBM round trips), and fuse the epilogue (do the bias/activation without re-reading the output).


7. Performance implications#

Shape matters enormously.

(4096, 4096) @ (4096, 4096)     → excellent: full tiles, high intensity
(1, 4096) @ (4096, 4096)        → GEMV: intensity ~1, memory bound, tensor cores mostly idle
(4096, 4096) @ (4096, 1)        → same problem
(129, 4096) @ (4096, 4096)      → tile quantization: 129 = 128+1, the second tile row is
                                   1/128 utilized. Can cost 2x.

Tile quantization is why you sometimes see a large jump in time when a dimension crosses a multiple of 128. Round batch sizes and hidden dimensions to friendly multiples where you can.

Wave quantization: if your grid has 133 thread blocks and the GPU has 132 SMs, you need two “waves,” the second nearly empty. Effective utilization ~50%.


8. Production implications#

  • Never write your own GEMM for production. cuBLAS/CUTLASS/cuBLASLt are the product of thousands of engineer-years. Write your own only to learn (Section VI.08) or to fuse something they can’t.
  • Do check which kernel is being selected. CUBLASLT_LOG_LEVEL, or Nsight Compute, will tell you. A shape that falls back to a generic kernel can be 3x slower.
  • Pad shapes to tensor-core-friendly multiples (8 for FP16, 16 for INT8/FP8) — this is why vocabulary sizes get padded to e.g. 128256 rather than 128000.
  • Watch for accidental GEMV. Batch 1 decode is inherently GEMV-shaped; that’s expected. But a GEMV where you expected a GEMM means a shape bug.

9. Common mistakes#

Getting the inner dimensions backwards. (M,K)@(K,N). If you see a shape error, print both shapes and read them left to right.

Assuming A@B == B@A. It doesn’t, and often can’t.

Ignoring association order. Can cost orders of magnitude.

Counting FLOPs as M·N·K instead of 2·M·N·K. Halves all your estimates.

Forgetting that the batch dimension is “free” in the matmul but not in memory. Bigger batch = same weight reads, more activation memory.


10. Hands-on exercise#

A. By hand. Compute these on paper, then verify:

1.  [2 1] @ [1 0]        2.  [1 2 3] @ [1]      3.  [1 0] @ [0 1]
    [0 3]   [2 1]                       [0]         [0 1]   [1 0]
                                        [2]

B. Write it three ways. Implement matmul in Go as (i) the textbook i, j, p triple loop, (ii) a sum of outer products, (iii) the i, p, j loop order. Verify all three agree. Time all three on 512×512. Explain the ratios.

C. Tiling. Implement a tiled matmul in Go and sweep tile size. Plot GFLOP/s vs tile size. Compare to the naive version, and to an optimized BLAS if you have one.

D. Tile quantization. Time (M, 4096) @ (4096, 4096) on GPU for M ∈ {127, 128, 129, 255, 256, 257}. Plot. Explain the jumps.

E. Association order. Time (A@B)@C vs A@(B@C) for the shapes in section 5. Confirm the 2000x.


11. Interview questions#

  1. State the shape rule and FLOP count for matrix multiplication.
  2. Explain matmul as a sum of outer products, and why hardware likes that view.
  3. What is tile quantization and how do you avoid it?
  4. Why is batch-1 decode a GEMV, and why does that matter?
  5. Given (1,4096) @ (4096,4096) @ (4096,4096), which association order do you choose and why?
  6. What does a GEMM kernel keep in registers, in shared memory, and in HBM?

12. Further reading#

  • [ESTABLISHED] Goto & van de Geijn, “Anatomy of High-Performance Matrix Multiplication”
  • [REFERENCE] NVIDIA “Matrix Multiplication Background User’s Guide” — tile/wave quantization
  • [REFERENCE] CUTLASS documentation and the CUTLASS GEMM API design docs
  • Next: 03 — Tensors, shapes, broadcasting

↑↓ navigate↵ openesc close