TOPIC #4Beginner 12 min read

Matrix Multiplication: Composition and the GEMM That Runs AI

AI
AI & ML Editorial
Report an issue
Key takeawayCore Concept Summary

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.

A Forward Pass Is a Chain of Matmuls
100%
Touchpad: Pinch to zoom • Drag to pan
Rendering visual architecture flowchart...

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

Insight

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

Insight

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 — A has k columns, B must 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 ≠ BA in general (and BA is often not even defined).
  • Order encodes the order you apply transformations — that's the whole story in section 4.
python— Matmul rules and the classic error in PyTorch: inner dims match (128=128), then swap them and it explodes
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 matmul

03.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.

code
   x ──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 AB applies B first, 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. AB and BA are 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.
Production Implementation in Big Tech
NVIDIA (cuBLAS / Tensor Cores)• Serving LLM training and inference workloads

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.

Exercise 1 of 30 answered
1

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?

Rate This Architecture ChapterFeedback & Rating

How clear and actionable was this distributed systems breakdown?

Related Concepts & Cross-References

Indexed from curriculum