Matrix Multiplication: Composition and the GEMM That Runs AI
Matrix multiplication is a stack of dot products: each output entry is one row meeting one column. The inner shapes must match, the order matters, and a whole neural-network forward pass is just a chain of these products — the single workload GPUs are built to run.
A Forward Pass Is a Chain of Matmuls
Each layer composes a linear map (matmul) with an element-wise nonlinearity. Inner dimensions must match; the output shape takes the outer dimensions.
01.The Problem: One Dot Product Is Not Enough
In the previous topic you learned the dot product: multiply two equal-length lists entry by entry and add up the result — one number that measures how aligned two vectors are.
A whole neural layer is the same question asked over and over. Imagine you have a batch of 32 inputs, and each of 512 neurons wants to score every input. That's 32 × 512 = 16,384 dot products, all at once.
So the question becomes
How do I do many dot products in one tidy operation, with rules for which shapes are even allowed?
The answer is matrix multiplication. A matrix (rank-2 tensor: a grid of numbers, from Topic 1) is just a convenient pile of rows and columns, and matmul is "do the dot product of every row with every column."
02.The Idea in Plain Words: Rows Meet Columns
The rule in one bold line
Each entry of the product AB is one dot product: row i of A against column j of B.
Formally, given A ∈ R^{m×k} and B ∈ R^{k×n}, the product C = AB ∈ R^{m×n} has entries:
Cᵢⱼ = Σₗ Aᵢₗ · Bₗⱼ
Read Σₗ as "add over the shared position." Two constraints to burn in:
- Inner shapes must match —
Ahas k columns,Bmust have k rows. Otherwise:matmul shapes error. - Output takes the outer shapes — (m×k)(k×n) → (m×n).
Algebra facts worth memorizing:
- Matrix multiplication is associative
(AB)C = A(BC)and distributive over addition. - It is not commutative:
AB ≠ BAin general (andBAis often not even defined). - Order encodes the order you apply transformations — that's the whole story in section 4.
import torch
A = torch.randn(32, 128) # batch of 32 vectors, dim 128
W = torch.randn(128, 512) # layer: maps 128 -> 512
C = A @ W # (32,512): inner 128 matches
print(C.shape) # torch.Size([32, 512])
try:
W @ A # (128,512) @ (32,128): 512 != 32
except RuntimeError as e:
print("shape mismatch") # the most common ML error there is
# Batched: attention scores for 8 heads at once
Q = torch.randn(8, 64, 256) # (heads, seq, head_dim)
K = torch.randn(8, 32, 256)
scores = Q @ K.transpose(1, 2) # (8, 64, 32) batched matmul03.A Tiny Worked Example, With Real Numbers
Smallest interesting case: one row times one column.
A = [1, 2] (shape 1×2) and B = [3, 4]ᵀ (shape 2×1, i.e. the column with 3 on top of 4).
- one dot product:
1·3 + 2·4 = 3 + 8 = 11 - result: the 1×1 matrix
[[11]]
Now add a second row to A: A = [[1, 2], [0, 5]] (2×2), with the same column B = [[3], [4]] (2×1).
- row 1 · column:
1·3 + 2·4 = 11 - row 2 · column:
0·3 + 5·4 = 20 - result:
[[11], [20]]— shape (2,1). Outer sizes 2 and 1 survived; the inner 2/2 vanished.
That vanishing is the rule in action. A had 2 columns, B had 2 rows → the shared 2 cancels, output takes the remaining outsides. Try to reverse it and B @ A gives shape (2,2) but completely different numbers — because AB ≠ BA.
04.Geometric Meaning: Composition of Transformations
Here is the deep "why". A matrix is a linear transformation — a rule that stretches, rotates, and shears space, fixed entirely by where it sends the basis vectors. Multiplying matrices composes those transformations: applying A after B is the single matrix AB.
codex ──W₁──► (W₁x) ──W₂──► W₂(W₁x) ──W₃──► W₃(W₂(W₁x)) input layer 1 out layer 2 out logits read right-to-left: the RIGHTMOST matrix acts FIRST
Consequences that pay off later:
- Stacking layers means composing maps; the deeper the stack, the more the product matrix can amplify or shrink certain directions (bridge to eigenvalues, Topic 7, and exploding gradients, Topic 12).
- Because
ABappliesBfirst, transformations read right to left:y = W₃(W₂(W₁x)). - If a product equals the identity matrix, the transformations undo each other (Topic 6, inverses).
The Mermaid diagram at the top draws the production version: input → matmul W₁ → ReLU → matmul W₂ → logits.
05.The Analogy: A Chain of Machines in a Factory
Carry one image through everything below: a factory where each machine reshapes the parts coming off the previous machine.
- A raw part = your input vector
x. - A machine = a matrix (a transformation: stretch, tilt, rotate).
- Bolt machine A onto machine B's output → you get one combined machine
BA. That's matrix multiplication: composing the factory line, not running it twice. - The part enters the rightmost machine first — the line reads right to left, like
y = W₃(W₂(W₁x)). - Two machines in the wrong order → broken part.
ABandBAare different factory lines. - A machine and its inverse, placed side by side, un-do each other: the product is the identity — the part comes out unchanged.
Now the punchline of the whole topic: a neural network is this factory. Each layer is "one matrix machine, then one small non-linear tweak" (ReLU), bolted into a long line. And the line only fits together if each joint's widths match — the inner-shape rule.
06.Why AI Cares: GEMM, the Hottest Operation in Hardware
Back to the factory: once every layer is a matmul, the machine that runs matmuls is the machine that runs AI.
A dense matmul of (m×k)(k×n) costs about 2·m·k·n FLOPs. Large language model training is dominated by a handful of gigantic matmuls per step — so much so that accelerator benchmarks are basically matmul benchmarks.
- The BLAS kernel GEMM (GEneral Matrix Multiply) and its batched/strided variants (cuBLAS, cuBLASLt, oneDNN) execute the bulk of FLOPs in every training and inference run.
- A 70B-parameter forward pass performs roughly 2×70×10⁹ FLOPs per token — a parade of (T×d)(d×4d) style products.
- Tensor cores exist because dot-product accumulation is the workload: they fuse multiply-add tiles for 16-bit and FP8 inputs (2024–2026: H100/B200-era mixed precision).
- Inference at batch size 1 is memory-bound: the matmuls are skinny (m = 1), so reading weights from HBM, not FLOPs, sets the speed — the reason for batching, and for techniques like continuous batching in vLLM and TensorRT-LLM.
One corollary from the factory picture, useful in interviews: stacking two linear layers with no activation between them is just one combined machine. Linear∘linear = linear — nonlinearity between the matmuls is what gives depth any value.
07.In Practice: Shape Hygiene and einsum
Real tensors are rank-N, so you multiply along chosen axes. NumPy/PyTorch's @ broadcasts over leading batch dimensions, and einsum makes axis intent explicit:
np.einsum("ij,jk->ik", A, B)— plain matmul.np.einsum("bhd,ehd->bhe", Q, K)— attention scores: batch × query-heads dot embedding-heads.
When code review asks "what are the shapes?", they are checking that inner dimensions agree and that batch axes broadcast. Practice reading products as "sum over the shared index" — the Σₗ in the formula from section 2, spoken aloud.
Architectural Trade-offs & Production Realities
Architectural Advantages
- Composes arbitrarily many linear maps into one product; associative, so you can re-associate for efficiency.
- O(2mnk) dense, regular access pattern — ideal for GPUs/TPUs and cache-friendly tiling.
- Expresses full layers, attention, convolutions (as im2col matmul), and PCA projections uniformly.
Trade-offs & Constraints
- Cost is cubic in width for square stacks; very large layers need sharding (tensor parallelism) or low-rank tricks.
- Non-commutativity and framework layout conventions (row vs column major) breed subtle bugs.
- Batch-1 inference is bandwidth-bound on weight loading, not compute-bound.
Hopper/Blackwell GPU architectures are tuned around one metric: mixed-precision GEMM throughput. A transformer training step is a sequence of (tokens×d_model)(d_model×d_ff) and back-propagation matmuls; kernel libraries tile these products to keep data in on-chip SRAM and stream weights from HBM at full bandwidth.
Staff+ Engineering Takeaways
- Entry (i,j) of AB is the dot product of row i of A and column j of B; inner shapes must match, output takes outer shapes.
- Matrix multiplication is composition of linear transformations, applied right-to-left, and is not commutative.
- Stacked linear layers collapse into a single matrix product — nonlinearity between them is what gives depth value.
- AI hardware revolves around GEMM: ~2mk n FLOPs per product, tensor cores are matmul accelerators.
- Batched matmul over leading axes (or einsum) is how attention scores are computed for all heads at once.
Topic Knowledge Check
Exercise 1 of 3 • Test your architectural comprehension.
X has shape (64, 512) and W has shape (d, 256). What must d be for X @ W to be defined, and what is the output shape?
How clear and actionable was this distributed systems breakdown?