primer.ml.hardware

The hardware underneath: chips, memory, links and number formats

Run: python -m primer.ml.hardware

New to the notation? primer.notation explains every symbol used here from zero. This lesson builds on the cost arithmetic of primer.ml.inference.

Level 1: The practitioner's guide

In one sentence. The hardware under a model is thousands of simple multipliers starved by slow memory, so what you rent or buy is decided by bytes (does the model fit, how fast can its weights be read, how fast can chips talk) and the number format you store those bytes in is the cheapest lever you have.

When you need it. You need this the day you have to pick a machine: a laptop for local experiments, a cloud GPU for a demo, a multi-GPU server for serving a 70-billion-parameter model, or a cluster for training. You also need it when a model runs far slower than its FLOP count suggests. The tell is a spec sheet you cannot read: teraFLOPS, HBM, NVLink, bf16, FP8, and no idea which number will bite. This lesson's imaginary datacenter GPU does about 10¹⁵ operations per second but reads only 3.35 TB/s, so any work doing fewer than about 299 operations per byte fetched leaves the multipliers idle; generating one token for one user does about 1. You do not need this lesson while you call a hosted API: the provider has done the sizing for you. You need it the moment the bill or the latency makes you consider doing it yourself.

Your options. From the least hardware to the most, at what each can hold and how its parts talk:

Option What it is What fits, roughly What it costs Where it lives
A hosted API Someone else's GPUs behind a per-token price Any model they offer, at any scale Money per token, no capacity planning, no control over the machine The provider
A laptop or consumer GPU One chip with a few to a few tens of gigabytes of memory, no fast links Small models, or larger ones quantized to 4 bits: an 8B model at 4 bits is 4 GB, a 70B is 35 GB (this lesson's memory math) Cheap and private; slow per token, one user at a time Your desk
One datacenter GPU 80 GB of HBM at 3.35 TB/s (NVIDIA's H100 SXM specification, and this lesson's constants) A 70B model only at 8 bits (70 GB, 10 GB of cache left) or 4 bits (35 GB, 45 GB left); a 16-bit 70B does not fit Rental by the hour; the whole card even when one user uses 0.3% of it A cloud instance or a rack
One machine, several GPUs on fast links Chips joined at hundreds of GB/s (NVLink is 900 GB/s on an H100 SXM; this lesson models 500) A model split across the GPUs, exchanging partial results inside every layer Several cards' rent; the fast links are what you are paying for A cloud instance or a rack
Many machines over a network Machines joined at tens of GB/s per GPU (this lesson models 50) Training runs and fleets: each machine holds a copy or a slice, and they talk once per step The most money and the most engineering; the network becomes the bottleneck A cluster

How to choose. Start from the model's size in bytes and the number format you are willing to run it in.

  • Compute the weights first: parameters times bytes per weight. If they fit in one GPU with room for the KV cache, stop there; one chip with no links is the simplest system you can operate.
  • If they do not fit, drop the format before adding chips: 8-bit weights halve the bytes and, on this lesson's numbers, cut the lower bound on decode time for a 70B model from 41.8 ms to 20.9 ms per token, and 4-bit to 10.4 ms. Check quality on your own tasks afterwards.
  • If they still do not fit, add GPUs inside one machine, where the links are fast enough to split a layer across chips.
  • Cross to many machines only for training or for a fleet, and design the split so that the chatty parallelism (tensor parallelism, talking inside every layer) stays within a machine and only once-per-step traffic crosses the network.
  • For training, pick bf16 for the multiplies and keep fp32 master weights: a gradient of 10⁻⁸ becomes exactly 0 in fp16 but survives in bf16, and a weight update of 0.001 vanishes in bf16 unless the master copy is fp32 (this lesson's format table). FP8 training is real and works on models up to 175B parameters with no hyperparameter changes (Micikevicius et al., FP8 Formats for Deep Learning), but it needs software that handles the scaling for you.
  • Whatever you pick, measure what fraction of peak FLOPS you reach. If it is 80%, you are at least 80% compute-bound; if it is a few percent, you are moving bytes, and more arithmetic will not help (Horace He, Making Deep Learning Go Brrrr).

What it costs. Memory is the price of admission and bandwidth is the speed limit. Reading 16 GB of weights once takes 4.78 ms from HBM on this lesson's GPU, 320 ms from the host's memory, 67 times slower: a model "offloaded" to CPU memory runs, but each token waits that much longer. Formats set both bills: halving the bits halves the bytes moved, doubles the operations per byte a tiled kernel achieves (63 becomes 126 in this lesson's 4096 × 4096 example), and shrinks the multiplier itself (an fp32 multiplier needs 576 cells of silicon, an fp8 one 16), which is why accelerators list roughly double the peak throughput at each halving of the format (the H100 lists 1,979 TFLOPS at bf16 and 3,958 at FP8, both with sparsity). What a format costs you in return is range or precision: fp16 tops out at 65,504, fp8 E4M3 at 448, and int4 holds only 15 levels, so small weights vanish without a per-row scale. Links cost time at scale: on this lesson's numbers an all-reduce of a 14 GB gradient across 8 GPUs takes 49 ms inside a machine and 490 ms across machines, against 690 ms of arithmetic per step, so the same run spends 7% of its time talking on fast links and 71% over a network. Power is part of the rent too: an H100 SXM is rated up to 700 W.

What breaks.

  • The model "fits" and then does not. Weights are the fixed cost; the KV cache grows with every token of every conversation, and activations and the serving software take several more gigabytes. Size for weights plus cache plus headroom, not weights alone.
  • A big GPU idles on a small job. A single user's decode reads every weight to do two operations with it. Without batching, most of the card you rent does nothing.
  • fp16 training silently zeros gradients: it loses precision below 6 × 10⁻⁵ and rounds anything under about 3 × 10⁻⁸ to zero (this lesson's format table). Use loss scaling, or use bf16, which trains to fp32 quality with no hyperparameter changes (Kalamkar et al.).
  • bf16 swallows small updates: 1 + 0.001 rounds back to 1. Keep master weights and long running sums in fp32.
  • fp8 overflows: 500 in E4M3 is not a number. The format needs per-tensor scaling that the training or serving library supplies; do not cast by hand.
  • Offloading to host memory makes a model fit at the price of tens of times slower steps. It is for experiments, not for serving.
  • Tensor parallelism across a network stalls in every layer. Keep it on the fast links inside a machine.

In the wild. NVIDIA's H100 specification gives the numbers this lesson rounds (80 GB at 3.35 TB/s, 900 GB/s NVLink). The formats each have a paper: Kalamkar et al. studied bf16 for training, and Micikevicius et al. proposed the two fp8 encodings, E4M3 and E5M2, and earlier the mixed-precision recipe (fp32 master weights, loss scaling) that PyTorch's automatic mixed precision, linked in Further reading, implements. The roofline model this lesson uses to decide memory-bound from compute-bound is Williams, Waterman and Patterson's, and FlashAttention (Dao et al.) is the best-known application of tiling to a model. Megatron-LM (Shoeybi et al.) is the tensor parallelism that lives on fast links, and Horovod (Sergeev and Del Balso) brought the ring all-reduce to deep learning. How to Scale Your Model, linked in Further reading, carries the same arithmetic through TPUs and GPUs to full training runs.

Go deeper. Level 2 builds each number here from nothing: a matrix multiply counted by hand, a memory hierarchy with its six levels, a tiled multiply whose traffic you can watch fall, a 16-bit float encoded bit by bit and every format's range and precision derived from its bit widths, a ring all-reduce simulated on four GPUs, and the serving-fit table computed from the formulas. If you only needed to choose a machine and a format, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

Every lesson so far has counted operations: so many multiplies per token, so many parameters. This lesson looks at the machine that performs them, because the machine explains things the maths alone never will: why a model that "needs" a tenth of a millisecond of arithmetic takes five milliseconds per token, why training runs are spread across thousands of chips in a particular way, and why everyone is shrinking numbers from 32 bits to 8.

Three facts carry the whole lesson:

  1. A GPU is thousands of simple arithmetic units doing the same step on different numbers. Neural networks are mostly matrix multiplies, which are exactly that kind of work.
  2. Arithmetic is cheap; moving data is expensive. The chip can multiply far faster than its memory can feed it, so speed is usually decided by how many bytes move, not how many operations run.
  3. Fewer bits per number helps everywhere at once: more numbers per byte moved, more numbers in memory, and smaller, more numerous multipliers on the chip.

Sections 5 and 6 then apply those facts to many GPUs working together, and to the question every deployment starts with: will the model fit?

1. Why GPUs: thousands of simple cooks

Everyday picture. A CPU is a few master chefs. Each can cook anything, improvise, and follow a recipe full of "if the sauce splits, do this instead". A GPU is a kitchen of thousands of line cooks who all do the same step at the same moment on different ingredients: "everyone, chop your carrot now." That kitchen is useless for inventing a menu and unbeatable at ten thousand identical salads. A neural network is ten thousand identical salads: nearly all of its work is matrix multiplication, the same multiply-and-add done billions of times on different numbers.

Tiny worked example. Multiply a 2 × 3 matrix by a 3 × 2 matrix:

Level 3: the formula and its symbols

$$ A = \begin{pmatrix} 1 & 2 & 3 \ 4 & 5 & 6 \end{pmatrix}, \quad B = \begin{pmatrix} 7 & 8 \ 9 & 10 \ 11 & 12 \end{pmatrix}, \quad AB = \begin{pmatrix} 58 & 64 \ 139 & 154 \end{pmatrix} $$

Symbols

Symbol Meaning here Shape
$A$ the left matrix: 2 rows of 3 numbers 2 × 3
$B$ the right matrix: 3 rows of 2 numbers 3 × 2
$AB$ the matrix multiply: the cell in row $i$, column $j$ is row $i$ of $A$ dotted with column $j$ of $B$ 2 × 2
$\begin{pmatrix}\ldots\end{pmatrix}$ a matrix written out, row by row

In words: "each cell of the answer is one row of A times one column of B, multiplied position by position and added up."

With the numbers: the top-left cell is 1·7 + 2·9 + 3·11 = 7 + 18 + 33 = 58; the bottom-right is 4·8 + 5·10 + 6·12 = 32 + 50 + 72 = 154. Each cell took 3 multiplies and 3 additions (each product added to a running total that starts at 0), so the 4 cells took 12 multiplies and 12 additions. The crucial detail: no cell needs any other cell's answer. Four cooks could take one cell each and finish at the same moment.

Level 3: in Python

In Python:

A = [[1, 2, 3], [4, 5, 6]]
B = [[7, 8], [9, 10], [11, 12]]
m, k, n = len(A), len(B), len(B[0])
# each cell: row i of A dotted with column j of B
[[sum(A[i][p] * B[p][j] for p in range(k)) for j in range(n)] for i in range(m)]  # → [[58, 64], [139, 154]]

That count generalises into the most useful formula in this lesson. A FLOP (floating-point operation) is one multiply or one add, and multiplying an m × k matrix by a k × n matrix costs:

Level 3: the formula and its symbols

$$ \text{FLOPs} = 2\,m\,n\,k $$

Symbols

Symbol Meaning here In the example
$m$ rows of the left matrix (and of the answer) 2
$k$ the shared inner size: columns of the left, rows of the right; the length of each dot product 3
$n$ columns of the right matrix (and of the answer) 2
$m\,n$ how many cells the answer has 4
$2$ one multiply plus one add per step of a dot product
FLOPs floating-point operations in total 24

In words: "every one of the m·n answer cells is a dot product of length k, and every step of a dot product is a multiply and an add."

With the numbers: 2 × 2 × 2 × 3 = 24 for the example. A 4096 × 4096 by 4096 × 4096 multiply, the size of one weight matrix in a mid-sized model, costs 2 × 4096³ ≈ 1.37 × 10¹¹ FLOPs. At the 10¹⁵ FLOPs per second of the imaginary datacenter GPU used throughout primer.ml.inference, that is about 0.14 milliseconds, if the chip could be kept busy.

Level 3: in Python

In Python:

m, n, k = 2, 2, 3
# 2 FLOPs per step, k steps per cell, m·n cells
2 * m * n * k  # → 24
flops = 2 * 4096 * 4096 * 4096
flops  # → 137438953472
# milliseconds at 10¹⁵ FLOPs per second
round(flops / 1e15 * 1000, 2)  # → 0.14

Hardware usually does "multiply, then add to a running total" as a single instruction, the fused multiply-add, which is why FLOPs come in pairs.

flowchart LR subgraph CPU["CPU: a few master chefs"] direction TB c1["core: big control unit,<br/>big cache, runs any code"] c2["core"] c3["core"] end subgraph GPU["GPU: thousands of line cooks"] direction TB g1["group of cores:<br/>one instruction, many numbers"] g2["group of cores"] g3["... about a hundred groups"] g4["matrix units: a small<br/>tile multiply per instruction"] end W["matrix multiply:<br/>m × n independent cells"] --> GPU BR["branchy code:<br/>if this then that"] --> CPU

Reading it: the two boxes spend their silicon differently. The CPU spends it on a few cores that each handle any instruction stream quickly, including code full of decisions. The GPU spends it on arithmetic: groups of cores that all execute the same instruction on different numbers, plus (on most modern accelerators) matrix units that multiply a small tile of numbers in one instruction. A matrix multiply is a perfect fit for the GPU, because its m × n cells are independent and identical in shape.

How far does the independence go? Suppose each core takes whole cells and performs one multiply-add per round:

Level 3: the formula and its symbols

$$ \text{rounds} = \left\lceil \frac{m\,n}{\text{cores}} \right\rceil \times k $$

Symbols

Symbol Meaning here In the example
$m\,n$ independent cells to compute 64 × 64 = 4,096
cores identical cores working at once 1 to 8,192
$\lceil x \rceil$ ceiling: round $x$ up to a whole number (half a cell of work still needs a round) $\lceil 0.5 \rceil = 1$
$k$ multiply-adds per cell, done one after another 64
rounds how long the whole multiply takes, in rounds

In words: "share the cells out evenly, round up, and each core then spends k rounds on each cell it was given."

With the numbers: a 64 × 64 by 64 × 64 multiply on 1 core takes 4,096 × 64 = 262,144 rounds; on 64 cores, 4,096; on 4,096 cores, 64. On 8,192 cores it is still 64: there are only 4,096 cells, so half the cores have nothing to do.

Level 3: in Python

In Python:

import math
m = n = k = 64
def rounds(cores):
    # ⌈m·n / cores⌉ cells per core, then k multiply-adds per cell
    return math.ceil(m * n / cores) * k
rounds(1), rounds(64), rounds(4096), rounds(8192)  # → (262144, 4096, 64, 64)

A 64 by 64 multiply speeds up in a straight line from 1 core to 4,096 cores, 262,144 rounds down to 64, then stays flat because there is no more independent work

Reading it: both axes are logarithmic. Doubling the cores halves the time, a straight line, until the cores match the number of independent cells (the dashed line at 4,096). Past that point the line goes flat: more cores cannot help a problem that has run out of independent work. A small matrix leaves a big GPU mostly idle.

In code: counted_matmul multiplies with plain loops and counts every multiply and add; matmul_flops is the 2·m·n·k formula; parallel_rounds is the rounds formula above.

Why it matters in practice. A GPU is fast only when it is given a lot of independent work at once: large matrices, and many sequences processed together. That is why serving systems batch requests together (primer.ml.inference) and why a small model answering one user at a time uses a sliver of the chip. It is also why neural networks look the way they do: architectures that turn into a few big matrix multiplies (the transformer, primer.ml.transformer) won partly because they suit this hardware, while step-by-step recurrences (primer.ml.cnn_rnn) do not.

2. The memory hierarchy: near is small, far is big

Everyday picture. Back in the kitchen. A cook's hands hold one or two things (the registers). The cutting board holds a few more (the on-chip SRAM, fast memory built into the chip itself). The fridge in the kitchen holds the day's ingredients (HBM, "high-bandwidth memory", the GPU's main memory). The storeroom down the hall is the CPU's memory (host memory). The warehouse across town is the disk. And other restaurants' pantries, reached by courier, are other machines over the network. Every step further out holds more and takes longer to reach. The cooks are fast; what slows the kitchen down is fetching.

Tiny worked example. Round, illustrative numbers for one datacenter GPU and the machine around it (orders of magnitude, not any product's specification):

Level Holds about Moves about Streaming 1 GB takes What lives there
registers 20 MB (across the chip) 100 TB/s 0.01 ms the numbers being multiplied this instant
on-chip SRAM 50 MB 20 TB/s 0.05 ms tiles of the current multiply (section 3)
HBM 80 GB 3.35 TB/s 0.30 ms weights, activations, the KV cache
host memory 1 TB 50 GB/s (over the link to the GPU) 20 ms data waiting to be loaded, offloaded state
local disk 10 TB 10 GB/s 100 ms datasets, checkpoints
network the whole cluster 50 GB/s per GPU 20 ms other GPUs' gradients, remote storage

Registers and SRAM never hold a whole gigabyte; the column shows their rate. Notice the jump from HBM to host memory: about 67 times slower. A GPU that has to reach past its own HBM is a cook walking to the storeroom for every carrot.

Level 3: the formula and its symbols

$$ t = \frac{\text{bytes}}{\text{bandwidth}} $$

Symbols

Symbol Meaning here Units
$t$ time to stream the data, ignoring the fixed delay before the first byte arrives (latency) seconds
bytes how much data moves bytes (1 GB = 10⁹)
bandwidth how many bytes per second the level can deliver bytes per second

In words: "the time to move data is its size divided by the speed of the pipe it moves through."

With the numbers: an 8-billion-parameter model at 2 bytes per parameter is 16 GB. Reading it once from HBM takes 16 × 10⁹ / 3.35 × 10¹² = 4.78 ms: exactly the lower bound on time per generated token in primer.ml.inference, because generating one token reads every weight once. From host memory it would take 320 ms.

Level 3: in Python

In Python:

HBM, host = 3.35e12, 50e9
weights = 16e9
# t = bytes / bandwidth, in milliseconds
round(weights / HBM * 1000, 2)  # → 4.78
round(weights / host * 1000)  # → 320
# how many times slower the storeroom is than the fridge
round(HBM / host)  # → 67
flowchart LR ALU["arithmetic units"] <--> R["registers<br/>~20 MB, ~100 TB/s"] R <--> S["on-chip SRAM<br/>~50 MB, ~20 TB/s"] S <--> H["HBM<br/>~80 GB, ~3.35 TB/s"] H <--> D["host memory<br/>~1 TB, ~50 GB/s"] D <--> K["local disk<br/>~10 TB, ~10 GB/s"] H <--> N["network: other machines<br/>~50 GB/s per GPU"]

Reading it: start at the arithmetic units on the left and walk outward. Every box to the right is bigger and slower. The arithmetic units can only work on numbers in registers, so every number used must travel the whole way in from wherever it lives. The network hangs off HBM because fast clusters let a GPU send and receive data straight from its own memory without a detour through the CPU.

Capacity grows about fifty-million-fold from registers to the network while bandwidth falls ten-thousand-fold from registers to disk

Reading it: the same six levels, top to bottom, on logarithmic axes. On the left, capacity climbs from megabytes to a petabyte. On the right, bandwidth falls from a hundred terabytes per second to ten gigabytes per second. The one place the order bends is the last two rows: a datacenter network is built to rival a local disk, so fetching from a nearby machine can be as fast as reading your own drive. Both are still around a hundred times slower than HBM.

In code: MEMORY_HIERARCHY lists the six levels with their capacity and bandwidth; transfer_seconds is the formula above.

Why it matters in practice. A model runs at full speed only if everything it touches every step (weights, activations, KV cache) lives in HBM. Spilling to host memory ("offloading") makes a model fit, at the price of each step waiting tens of times longer. And because HBM itself is slow compared with the arithmetic, the fastest code is the code that makes each trip to HBM count, which is the next section.

3. Arithmetic intensity: why data movement dominates

Everyday picture. A sandwich shop. If the cook walks to the storeroom for each slice of bread for each sandwich, the cook spends the day walking. If the cook carries a tray of bread and a tray of fillings to the bench and makes a batch of sandwiches from them, every trip feeds many sandwiches. Same sandwiches (FLOPs), far fewer trips (bytes). The ratio of the two is the arithmetic intensity: operations done per byte fetched.

Our imaginary GPU does 10¹⁵ FLOPs per second but reads only 3.35 × 10¹² bytes per second from HBM, so it breaks even at about 299 FLOPs per byte (the ridge point of the roofline in primer.ml.inference). Any work doing fewer operations than that for each byte it fetches leaves the arithmetic units waiting on memory.

Tiny worked example. Multiply two 4 × 4 matrices: 2 × 4³ = 128 FLOPs. Count the numbers fetched from slow memory (HBM) into fast memory (SRAM):

Strategy Numbers read Numbers written FLOPs per number moved
no reuse: each cell fetches its own row of A and column of B 16 × (4 + 4) = 128 16 128 / 144 = 0.89
2 × 2 tiles: load a tile of A and a tile of B, use each number twice 64 16 128 / 80 = 1.6
one 4 × 4 tile: load everything once 32 16 128 / 48 = 2.7

The arithmetic is identical in all three rows. Only the traffic changes. This trick is tiling: bring a small block of each matrix into fast memory and do every multiplication that block takes part in before throwing it away.

Level 3: the formula and its symbols

$$ \text{reads} = \frac{2\,n^3}{T} \qquad I = \frac{2\,n^3}{b\left(\dfrac{2\,n^3}{T} + n^2\right)} \approx \frac{T}{b} $$

Symbols

Symbol Meaning here In the example
$n$ both matrices are $n \times n$ 4, then 4,096
$T$ tile width: fast memory works on $T \times T$ blocks 1, 2, 4, then 128
$2\,n^3$ the FLOPs of the whole multiply (section 1, with $m = n = k$) 128
reads numbers fetched from slow memory 128, 64, 32
$n^2$ numbers written back: each answer cell once 16
$b$ bytes per number: 2 at 16-bit, 1 at 8-bit 2
$I$ arithmetic intensity: FLOPs per byte moved FLOPs/byte
$\approx$ "roughly", once $n$ is much bigger than $T$ and the writes are negligible

In words: "every number fetched is used T times, so the traffic falls in proportion to the tile width, and the intensity rises in proportion to it."

With the numbers: for n = 4, the reads are 2 × 64 / T = 128, 64 and 32 for tiles of 1, 2 and 4, as the table says. For two 4096 × 4096 matrices in 16-bit with 128-wide tiles, I ≈ 128 / 2 = 64, and 63.0 once the writes are counted. In 8-bit the same tiles give 126: halving the bytes per number doubles the intensity.

Level 3: in Python

In Python:

n = 4
# reads = 2n³ / T, for tiles of 1, 2 and 4
[2 * n**3 // T for T in (1, 2, 4)]  # → [128, 64, 32]
def intensity(n, T, b):
    flops = 2 * n**3
    # bytes moved: every read, plus one write per answer cell
    moved = b * (2 * n**3 / T + n**2)
    return flops / moved
round(intensity(4096, 128, 2), 1)  # → 63.0
round(intensity(4096, 128, 1), 1)  # → 126.0
flowchart LR subgraph HBM["HBM: big, slow"] A["A, in T × T tiles"] B["B, in T × T tiles"] C["C, the answer"] end subgraph SRAM["on-chip SRAM: small, fast"] a["one tile of A"] b["one tile of B"] acc["running total for<br/>one T × T tile of C"] end A -- "load" --> a B -- "load" --> b a --> mm["multiply-add:<br/>T³ steps, no traffic"] b --> mm mm --> acc mm -. "next pair of tiles along k" .-> A acc -- "write once, at the end" --> C

Reading it: the left box is slow memory, the right box fast memory. For one tile of the answer, the kernel repeatedly loads one tile of A and one tile of B, does all T³ multiply-adds between them without touching slow memory, and adds the results into a running total that stays on chip. Only when the whole row of tiles has been consumed does it write the finished answer tile back, once. Fast memory never holds more than three tiles.

Measured reads fall from 65,536 to 2,048 as the tile grows from 1 to 32, on the 2n-cubed-over-T line, and intensity at n = 4096 rises with the tile, reaching the 299 break-even only near T = 650 in 16-bit or T = 310 in 8-bit

Reading it: on the left, dots are reads counted by actually running the tiled multiply on 32 × 32 matrices, and the line is the formula 2n³/T; they agree exactly, and each doubling of the tile halves the traffic. On the right, the intensity of a 4096 × 4096 multiply climbs with the tile width, and the dashed line is the chip's break-even point of 299. In 16-bit the tiles need to be about 650 wide to cross it. Three 650 × 650 tiles in 16-bit take about 2.5 MB, more fast memory than one group of cores has to itself, which is why real kernels tile at several levels at once (tiles in registers inside tiles in SRAM) and why lower-precision numbers, which shift the whole curve up, are so attractive.

The same idea runs through the rest of this primer:

  • FlashAttention (primer.ml.attention) tiles attention: blocks of queries, keys and values are loaded into SRAM, and the n × n score matrix is never written to HBM at all. Same answer, a fraction of the traffic.
  • Kernel fusion: adding a bias or applying an activation does about one FLOP per number it reads, hopelessly below 299, so these steps are done inside the matrix-multiply kernel while the tile is still on chip.
  • Decode (primer.ml.inference): generating one token for one user reads every weight to do just 2 FLOPs with it, an intensity of about 1. Batching users together is tiling across requests: one read of a weight serves every sequence in the batch.

In code: tiled_matmul runs the tiled multiply and returns a Traffic count of reads, writes, FLOPs and peak fast-memory use; matmul_reads and matmul_intensity are the formulas; primer.ml.inference.ridge_point is the break-even.

Why it matters in practice. Before asking how many FLOPs a piece of work needs, ask how many bytes it moves and how often each byte is reused. Most large speed-ups in modern AI systems (FlashAttention, fused kernels, batching, quantization) change the bytes, not the FLOPs.

4. Number formats: how many bits each number gets

Everyday picture. Scientific notation on a form with a fixed number of boxes: 6.02 × 10²³. One box holds the sign. A few boxes hold the power of ten, which sets how big or small the number can be: its range. The rest hold the digits, which set how finely it is measured: its precision. With a fixed number of boxes, moving a box from the digits to the power buys range and costs precision. Computers do the same with bits and powers of two, and call the parts the sign, the exponent and the mantissa (the stored digits).

Tiny worked example. Store −6.5 in bf16 ("brain float 16": 1 sign bit, 8 exponent bits, 7 mantissa bits).

  1. Sign: negative, so the sign bit is 1.
  2. Power of two: the largest power of two not above 6.5 is 4 = 2², so 6.5 = 1.625 × 2².
  3. Exponent: stored with a bias of 127 added, so that negative powers need no sign of their own: 2 + 127 = 129 = 10000001 in binary.
  4. Mantissa: the leading 1 of 1.625 is always there, so it is not stored. The fraction 0.625 = ½ + ⅛ is 0.101 in binary, padded to seven bits: 1010000.
  5. The 16 bits: 1 10000001 1010000.

Most numbers are not so lucky. 0.1 has no finite binary expansion, so it is rounded to the nearest value each format can hold: 0.10000000149 in fp32, 0.1000977 in bf16, 0.0999756 in fp16 and 0.1015625 in 8-bit E4M3.

Level 3: the formula and its symbols

$$ x = (-1)^{s} \times 2^{\,e - \text{bias}} \times \left(1 + \frac{f}{2^{M}}\right), \qquad \text{bias} = 2^{E-1} - 1 $$

Symbols

Symbol Meaning here In the example
$s$ the sign bit: 0 positive, 1 negative 1
$(-1)^s$ −1 multiplied by itself $s$ times: +1 when $s = 0$, −1 when $s = 1$ −1
$E$ how many exponent bits the format has 8
$e$ the stored exponent, read as an ordinary whole number 129
bias the offset subtracted from $e$, so stored values 1…254 stand for powers −126…127 127
$M$ how many mantissa bits the format has 7
$f$ the stored mantissa, read as a whole number from 0 to $2^M - 1$ 1010000 = 80
$1 + f/2^M$ the significand: the hidden leading 1 plus the stored fraction, between 1 and 2 1 + 80/128 = 1.625
$x$ the number the bits stand for −6.5

In words: "the sign says plus or minus, the exponent says which power of two to scale by, and the mantissa says how far between that power and the next one the number sits."

With the numbers: (−1)¹ × 2^(129 − 127) × (1 + 80/128) = −1 × 4 × 1.625 = −6.5.

Level 3: in Python

In Python:

x = -6.5
M, bias = 7, 127
# s: 1 for a negative number
s = 1 if x < 0 else 0
# e: the power of two below |x| is 2², stored with the bias added
e = 2 + bias
e, format(e, "08b")  # → (129, '10000001')
# f: the fraction after the hidden 1 of |x| / 2², as a 7-bit whole number
f = round((abs(x) / 2**2 - 1) * 2**M)
f, format(f, "07b")  # → (80, '1010000')
# decode: (-1)^s × 2^(e - bias) × (1 + f / 2^M)
(-1)**s * 2**(e - bias) * (1 + f / 2**M)  # → -6.5

Two corners of the formula matter in practice. When the stored exponent is 0, the hidden 1 is dropped and the number is a subnormal: it lets values fade gradually towards zero instead of dropping off a cliff, at the cost of fewer significant bits. And IEEE-style formats reserve the all-ones exponent for infinity and NaN ("not a number"), which is where overflowing values go.

flowchart LR X["x = −6.5"] --> S["sign: negative<br/>s = 1"] X --> P["largest power of two<br/>not above 6.5: 2² = 4"] P --> E["exponent: 2 + bias 127<br/>e = 129 = 10000001"] P --> F["6.5 / 4 = 1.625<br/>drop the leading 1: .625"] F --> R["round .625 to 7 bits<br/>f = 1010000"] S --> B["1 | 10000001 | 1010000"] E --> B R --> B

Reading it: a number enters on the left and splits three ways. The sign is read off directly. The exponent comes from finding which pair of powers of two the number sits between. The mantissa is where the number sits within that pair, rounded to however many bits the format allows: this rounding box is the only place information is lost, and it is where every format differs.

Bit layouts drawn to scale: fp32 has 1 sign, 8 exponent and 23 mantissa bits; bf16 keeps the 8 exponent bits and cuts the mantissa to 7; fp16 has 5 and 10; the two fp8 formats have 5 and 2, or 4 and 3; int8 and int4 are plain integers

Reading it: each row is one format drawn to scale, one cell per bit, with the sign in grey, the exponent in orange and the mantissa in blue. Compare bf16 and fp16: the same 16 bits, split differently. bf16 is literally the top half of fp32, keeping all 8 exponent bits (fp32's whole range) and giving up precision; fp16 keeps more precision and gives up range. The integer formats at the bottom have no exponent at all: every value is a whole number, turned into a weight by one shared scale (primer.ml.inference builds int8 and int4 quantization from scratch).

Format Bits (sign, exponent, mantissa) Largest Smallest at full precision Gap just above 1 Typical use
fp32 1, 8, 23 3.4 × 10³⁸ 1.2 × 10⁻³⁸ 1.2 × 10⁻⁷ master weights, optimizer state, running sums
bf16 1, 8, 7 3.4 × 10³⁸ 1.2 × 10⁻³⁸ 1/128 ≈ 0.0078 training and inference matrix multiplies
fp16 1, 5, 10 65,504 6.1 × 10⁻⁵ 1/1024 ≈ 0.00098 inference; training with loss scaling
fp8 E5M2 1, 5, 2 57,344 6.1 × 10⁻⁵ 0.25 gradients in 8-bit training
fp8 E4M3 1, 4, 3 448 0.0156 0.125 weights and activations in 8-bit
int8 8-bit integer 127 × scale evenly spaced none: a fixed step quantized weights
int4 4-bit integer 7 × scale evenly spaced none: a fixed step quantized weights

E4M3 bends the IEEE rules: it has no infinity, and spends that exponent on ordinary numbers instead, which is how 8 bits reach 448 rather than 240.

Relative spacing between neighbouring values: each float format is a flat band across its range (fp32 near 1e-7, bf16 near 1e-2, fp8 near 0.1), rising at its small end and stopping at its largest value, while int8 and int4 spacing rises steadily as numbers shrink

Reading it: the x-axis is the size of the number being stored; the y-axis is the gap to the next storable value, as a fraction of the number (lower is more precise); both are logarithmic. Each float format is a flat band: its relative precision is the same for tiny and huge numbers, which is the whole point of an exponent. The band's width is the range: it stops on the right at the largest value (65,504 for fp16, 448 for E4M3) and rises on the left where subnormals run out of bits. fp32 and bf16 are flat across the whole plot and far beyond it. The integer formats, here with a scale that maps 1.0 to the top code, are straight rising lines: their step is the same size everywhere, so small numbers get coarse, which is why quantized models use one scale per row or block of weights.

What the picture means for training (primer.ml.pretraining covers mixed precision in full):

  • Range failures. A gradient of 10⁻⁸ becomes exactly 0 in fp16 (below its smallest subnormal) but survives in bf16 as 1.0012 × 10⁻⁸. fp16 training therefore multiplies the loss by a large constant (loss scaling) to lift gradients into range; bf16 training does not need to.
  • Precision failures. In bf16, 1 + 0.001 rounds back to exactly 1: a small weight update simply vanishes. So training keeps a master copy of the weights in fp32, adds its sums up in fp32, and uses 16-bit (or 8-bit) only for the big multiplies. That split is mixed precision.

Why smaller formats multiply throughput

Fewer bits per number pays three times:

  1. Bytes. Half the bytes per number moves twice the numbers per second through every level of section 2, and fits twice the parameters in HBM. Memory-bound work speeds up directly: the lower bound on time per token for a 70-billion-parameter model (primer.ml.inference) is 41.8 ms at 16-bit, 20.9 ms at 8-bit and 10.4 ms at 4-bit.
  2. Intensity. The same tile does twice the FLOPs per byte (section 3: 63 becomes 126), pushing more work past the break-even point.
  3. Silicon. A multiplier is the expensive part of the chip, and its size grows with the square of the number of significand bits.

Everyday picture. Long multiplication by hand: multiplying two 3-digit numbers means writing a 3 × 3 grid of single-digit products; two 6-digit numbers need a 6 × 6 grid, four times the work. A chip's multiplier is that grid built in wires.

Level 3: the formula and its symbols

$$ \text{cells} = p^{2}, \qquad p = M + 1 $$

Symbols

Symbol Meaning here In the example
$M$ stored mantissa bits 23, 10, 7, 3
$p$ significand bits: the stored mantissa plus the hidden leading 1 24, 11, 8, 4
cells one-bit products in a schoolbook (array) multiplier: one per pair of bits

In words: "multiplying two p-bit significands needs one small cell for every pair of bits, so p times p cells; the exponents only need adding, which is cheap."

With the numbers: fp32 needs 24² = 576 cells, fp16 11² = 121, bf16 8² = 64, and fp8 E4M3 4² = 16: one fp32 multiplier's worth of silicon holds many 8-bit ones. Accelerators typically list roughly double the peak operations per second each time the format halves.

Level 3: in Python

In Python:

# stored mantissa bits for fp32, fp16, bf16 and fp8 E4M3
mantissa_bits = [23, 10, 7, 3]
# p = M + 1, and cells = p²
[(M + 1) ** 2 for M in mantissa_bits]  # → [576, 121, 64, 16]

In code: FloatFormat describes a format by its exponent and mantissa bits, with FloatFormat.max_value, FloatFormat.min_normal, FloatFormat.min_subnormal and FloatFormat.epsilon derived from them; FP32, BF16, FP16, FP8_E5M2 and FP8_E4M3 are the five formats; encode rounds a number into its three fields with plain arithmetic, decode turns fields back into a number, round_to does both, and bit_string prints the bits; multiplier_cells is the p² formula. The tests check round_to against NumPy's own float16 and float32.

Why it matters in practice. Picking a number format is picking a point on the range-versus-precision curve for each kind of number in a model. Weights and activations tolerate coarse formats; gradients need range; the running sums of long dot products need precision. Modern training and serving use a different format for each, and the savings are among the largest in the field.

5. Many GPUs: the cost of talking

Everyday picture. A group project. Four people each work through a quarter of the exercises, and then must agree on one combined answer sheet. Sitting at the same table they can compare notes in seconds; living in different cities they must post letters. The more often a team must compare notes, the more it matters who sits at the same table.

Large models are trained on many GPUs because no single one has the memory or the speed. The simplest split is data parallelism: every GPU holds a copy of the model and works on different examples, and after each step their gradients must be added up, so every copy takes the same step. That "add up, and give everyone the total" operation is an all-reduce.

Tiny worked example: a ring all-reduce. Four GPUs each hold a gradient of 8 numbers. Each splits its gradient into 4 chunks of 2 numbers and they sit in a ring, each passing to its right-hand neighbour.

  1. Reduce-scatter, 3 steps: each GPU sends one chunk to its neighbour, which adds it to its own copy of that chunk. After 3 steps, each GPU holds one chunk that contains the sum from all four.
  2. All-gather, 3 more steps: the finished chunks travel round the ring, overwriting the stale copies.

Each GPU sent 6 chunks of 2 numbers, 12 numbers, which is 2 × 3/4 × 8. The remarkable part: with 400 GPUs instead of 4, each would still send just under twice its gradient. The work per GPU barely grows.

flowchart LR G0["GPU 0<br/>chunks a b c d"] -- "one chunk per step" --> G1["GPU 1<br/>chunks a b c d"] G1 -- "one chunk per step" --> G2["GPU 2<br/>chunks a b c d"] G2 -- "one chunk per step" --> G3["GPU 3<br/>chunks a b c d"] G3 -- "one chunk per step" --> G0

Reading it: four GPUs in a ring, each only ever talking to its right-hand neighbour. Every link is busy at every step, each carrying a different chunk, so no link sits idle and no single GPU becomes a bottleneck. After 2 × (4 − 1) = 6 steps, every GPU has every chunk summed over all four.

Level 3: the formula and its symbols

$$ t_{\text{all-reduce}} = \frac{2\,(N-1)}{N} \cdot \frac{S}{B} \qquad t_{\text{compute}} = \frac{6\,P\,D}{F} $$

Symbols

Symbol Meaning here In the example
$N$ GPUs taking part 8
$S$ bytes of gradient each GPU holds 7 × 10⁹ parameters × 2 bytes = 14 GB
$B$ each GPU's link bandwidth 500 GB/s in one machine, 50 GB/s between machines (illustrative)
$\frac{2(N-1)}{N}$ the fraction of its gradient each GPU sends in a ring all-reduce: just under 2 1.75
$P$ parameters in the model 7 × 10⁹
$D$ tokens each GPU processes per step 16,384
$6$ FLOPs per parameter per token in training: 2 forward, 4 backward
$F$ the GPU's arithmetic speed 10¹⁵ FLOPs per second

In words: "talking takes just under twice the gradient's size divided by the link speed, however many GPUs there are; computing takes six operations per parameter per token, divided by the chip's speed."

With the numbers: inside one machine, 1.75 × 14 × 10⁹ / 500 × 10⁹ = 49 ms. Across machines, 490 ms. The arithmetic for the step is 6 × 7 × 10⁹ × 16,384 / 10¹⁵ = 0.69 s. Inside a machine, talking costs 7% of the computing time; across machines, 71%.

Level 3: in Python

In Python:

N, S = 8, 14e9
in_machine, between_machines = 500e9, 50e9
# 2(N - 1)/N × S / B, in milliseconds
round(2 * (N - 1) / N * S / in_machine * 1000)  # → 49
round(2 * (N - 1) / N * S / between_machines * 1000)  # → 490
P, D, F = 7e9, 16_384, 1e15
# 6·P·D / F, in seconds
round(6 * P * D / F, 2)  # → 0.69

All-reduce time levels off as GPUs are added: about 56 ms over fast in-machine links and about 560 ms over the network, against 690 ms of arithmetic per step

Reading it: the x-axis is the number of GPUs (log scale); the y-axis is seconds per training step. The flat grey line is the arithmetic every GPU does per step. The two rising curves are the all-reduce over each kind of link, and both flatten out almost at once: that is the ring's 2(N − 1)/N approaching 2. The gap between the curves is the whole story: the same gradient costs ten times more over the network than over in-machine links, bringing communication close to the cost of the arithmetic itself. Real systems hide part of it by sending the gradients of later layers while earlier layers are still computing theirs.

flowchart TB subgraph M1["machine 1: fast links, ~500 GB/s"] a1["GPU"] <--> a2["GPU"] <--> a3["GPU"] <--> a4["GPU"] end subgraph M2["machine 2: fast links, ~500 GB/s"] b1["GPU"] <--> b2["GPU"] <--> b3["GPU"] <--> b4["GPU"] end M1 <-- "network, ~50 GB/s per GPU:<br/>data parallelism, once per step" --> M2 TP["tensor parallelism:<br/>talks inside every layer"] -.-> M1 TP -.-> M2

Reading it: two machines, each with a few GPUs joined by fast links, and a slower network between them. The chattiest kind of splitting, tensor parallelism (cutting each matrix multiply across GPUs, which must swap partial results inside every layer), is kept within one machine's fast links. Kinds that talk rarely, like data parallelism (one all-reduce per step) and pipeline parallelism (passing activations only between neighbouring stages of layers), span the slower network. The layout of the wires decides the layout of the work. primer.ml.pretraining builds these strategies in full.

In code: ring_all_reduce simulates the ring step by step and counts the numbers each GPU sends; all_reduce_seconds and training_step_seconds are the two formulas; IN_MACHINE_LINK and BETWEEN_MACHINES_LINK are the illustrative link speeds.

Why it matters in practice. At scale, a cluster's network matters as much as its chips. Parallelism strategies are chosen by matching how often each kind talks to how fast each link is, and a training run that ignores the map spends its budget waiting.

6. Will it fit? Memory math for serving

Everyday picture. A bookshelf of fixed width. The encyclopedia (the model's weights) must go on it, whole. Whatever space is left holds one notebook per customer being served (their KV cache), and a notebook grows with every word of the conversation. Thinner volumes, printed in a smaller number format, leave room for more notebooks.

Tiny worked example. A model shaped like Llama 3 70B (80 layers, 8 key/value heads of 128 dimensions each) on one 80 GB GPU. Its KV cache costs 2 × 80 × 8 × 128 × 2 = 327,680 bytes per token at 16-bit (primer.ml.inference derives this).

Weights Weight memory Left for the cache Tokens of 16-bit cache Tokens of 8-bit cache
16-bit 140 GB none: does not fit 0 0
8-bit (fp8) 70 GB 10 GB 30,517 61,035
4-bit 35 GB 45 GB 137,329 274,658

This ignores activations and the serving software's own working memory, which take several more gigabytes in practice.

Level 3: the formula and its symbols

$$ P \, b_w + C \times 2\,L\,H_{kv}\,d_h\,b_{kv} \le M_{\text{GPU}} $$

Symbols

Symbol Meaning here In the example
$P$ parameters 7 × 10¹⁰
$b_w$ bytes per weight: bits ÷ 8 2, 1 or 0.5
$C$ tokens held in the KV cache, summed over every request being served the unknown
$2$ one key and one value per token
$L$ layers, each with its own cache 80
$H_{kv}$ key/value heads 8
$d_h$ dimensions per head 128
$b_{kv}$ bytes per cached number 2 or 1
$M_{\text{GPU}}$ the GPU's memory 80 × 10⁹ bytes
$\le$ "must be at most"

In words: "the weights, plus every cached token's keys and values across every layer, must fit in the GPU's memory."

With the numbers: with fp8 weights, 7 × 10¹⁰ × 1 = 70 GB, leaving 10 GB. 10 × 10⁹ / 327,680 = 30,517 tokens of 16-bit cache, or 61,035 with the cache in fp8 too: one long conversation, or a dozen short ones.

Level 3: in Python

In Python:

P, M_gpu = 70e9, 80e9
L, H_kv, d_h = 80, 8, 128
def kv_per_token(b_kv):
    # a key and a value, per layer, per KV head
    return 2 * L * H_kv * d_h * b_kv
kv_per_token(2)  # → 327680
# weight GB at 16, 8 and 4 bits
[P * bits / 8 / 1e9 for bits in (16, 8, 4)]  # → [140.0, 70.0, 35.0]
# tokens of cache beside fp8 weights, with a 16-bit cache and then an fp8 one
int((M_gpu - P * 1) // kv_per_token(2)), int((M_gpu - P * 1) // kv_per_token(1))  # → (30517, 61035)
flowchart LR P["parameters × bytes per weight"] --> W["weights"] T["tokens × bytes per token"] --> K["KV cache"] W --> Q{"weights + cache<br/>fit in GPU memory?"} K --> Q Q -- yes --> Y["serve: spare room means<br/>more users or longer contexts"] Q -- no --> N["smaller formats, fewer KV heads,<br/>shorter contexts, or more GPUs"]

Reading it: two budgets feed one question. The weights are a fixed cost, set by the parameter count and the number format. The cache grows with every token of every active conversation. If they fit together, whatever is left over becomes capacity: more concurrent users, longer contexts. If they don't, every lever on the right shrinks one of the two inputs, or splits the model across GPUs joined by the fast links of section 5.

Stacked bars against an 80 GB line: 16-bit weights alone reach 140 GB and overflow; fp8 weights take 70 GB and leave 10 GB, about 30,000 tokens of cache; 4-bit weights take 35 GB and leave 45 GB, about 137,000 tokens

Reading it: each bar is one choice of weight format for the same 70-billion-parameter model. The dark part is the weights; the light part is the room left for the KV cache, labelled with how many 16-bit tokens fit in it. The dashed line is the GPU's 80 GB. At 16-bit the weights alone break through the line. Each halving of the weight format frees memory that turns directly into tokens of cache, which is to say into users.

In code: max_cache_tokens subtracts the weights and divides what is left by the per-token cache, using primer.ml.inference.weight_bytes and primer.ml.inference.kv_cache_bytes_per_token; LLAMA3_70B_SHAPE holds the example model's shape.

Why it matters in practice. This arithmetic comes first in every deployment: the number format sets whether a model fits, how many users share a GPU, and (because decode reads every weight per token) how fast each of them sees words appear. primer.ml.inference carries it on into batching, speculative decoding and prompt caching.

In 20 seconds

  • GPUs are thousands of simple cores doing the same step on different numbers. Matrix multiplies (2·m·n·k FLOPs, every cell independent) are the perfect workload, as long as there is enough independent work.
  • Memory hierarchy: registers, on-chip SRAM, HBM, host memory, disk, network. Each step out is bigger and slower; the GPU runs at full speed only on data in HBM or closer.
  • Arithmetic intensity: the chip does hundreds of operations in the time it fetches one byte, so speed is set by bytes moved. Tiling reuses each fetched number T times and cuts traffic T-fold; FlashAttention, kernel fusion and batching are the same idea.
  • Number formats: sign, exponent (range) and mantissa (precision). bf16 keeps fp32's range with less precision; fp16 the reverse; fp8 and int8/int4 trade more. Fewer bits means fewer bytes, higher intensity and smaller multipliers, so throughput roughly doubles per halving.
  • Many GPUs: a ring all-reduce costs about 2 × gradient size ÷ link speed, whatever the GPU count. Chatty parallelism stays on fast in-machine links; the rest crosses the network.
  • Serving: weights plus KV cache must fit; the number format decides both what fits and how fast each token comes.

Self-test questions

Why are GPUs, rather than CPUs, used to train and run neural networks? Nearly all of a network's work is matrix multiplication, where every output cell is an independent dot product. A GPU spends its silicon on thousands of simple arithmetic units (plus matrix units) that apply the same instruction to different numbers, so it can compute thousands of cells at once. A CPU spends its silicon on a few flexible cores that are better at branchy, sequential code.

How many FLOPs does multiplying a 1,000 × 2,000 matrix by a 2,000 × 500 matrix take, and why that formula? 2 × 1,000 × 500 × 2,000 = 2 × 10⁹. There are m·n = 500,000 output cells, each a dot product of length k = 2,000, and each step of a dot product is one multiply and one add.

The chip can do 10¹⁵ FLOPs per second but a model runs far slower. What is usually the bottleneck, and how do you tell? Memory bandwidth. Compare the work's arithmetic intensity (FLOPs per byte moved from HBM) with the chip's break-even ratio, peak FLOPs divided by bandwidth (about 299 here). Below it, the arithmetic units wait on memory; generating one token for one user has an intensity near 1, so it is deeply memory-bound.

How does tiling a matrix multiply reduce memory traffic, and what limits it? Each block of numbers is loaded into fast on-chip memory once and used for every multiplication it takes part in, T times, instead of being fetched again for each one. Reads fall from 2n³ to 2n³/T. The limit is fast-memory size: three T × T tiles must fit at once, so real kernels tile at several levels of the hierarchy.

What is the difference between bf16 and fp16, and why do many training runs prefer bf16? Both have 16 bits. bf16 has fp32's 8 exponent bits and 7 mantissa bits: the same range as fp32, less precision. fp16 has 5 exponent bits and 10 mantissa bits: more precision, but a range that tops out at 65,504 and rounds gradients smaller than about 3 × 10⁻⁸ to zero. bf16 avoids the need for loss scaling; its coarse precision is handled by keeping fp32 master weights and fp32 sums.

Why does halving the bits per number roughly double throughput? Three reasons: half the bytes to move, so memory-bound work runs twice as fast and twice the parameters fit; twice the FLOPs per byte for the same tile, so more work clears the break-even point; and multipliers whose area grows with the square of the significand bits, so many more small multipliers fit in the same silicon.

Why is tensor parallelism usually kept inside one machine while data parallelism spans many? Tensor parallelism splits every matrix multiply, so GPUs must exchange partial results inside every layer, many times per step: it needs the fast in-machine links. Data parallelism communicates once per step (an all-reduce of the gradients), which a ring spreads so each GPU sends only about twice its gradient, and which can overlap with the backward pass, so it tolerates the slower network.

Will a 70-billion-parameter model serve from one 80 GB GPU? Not in 16-bit: the weights alone are 140 GB. In 8-bit the weights take 70 GB, leaving about 10 GB, around 30,000 tokens of 16-bit KV cache for a model with 80 layers and 8 KV heads of 128 dimensions. In 4-bit, 45 GB is left, about 137,000 tokens. Leave headroom for activations and the serving software.

The papers behind this lesson

  • Williams, Waterman and Patterson, Roofline: An Insightful Visual Performance Model for Multicore Architectures (2009): https://doi.org/10.1145/1498765.1498785. Introduced the roofline: judge a kernel by its arithmetic intensity against the machine's balance of compute and bandwidth.
  • Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (2022): https://arxiv.org/abs/2205.14135. Applied tiling to attention so the score matrix never reaches HBM, showing that counting memory traffic rather than FLOPs is what makes attention fast. Annotated companion
  • Micikevicius et al., Mixed Precision Training (2017): https://arxiv.org/abs/1710.03740. Showed that networks train in 16-bit floats with fp32 master weights, fp32 accumulation and loss scaling.
  • Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training (2019): https://arxiv.org/abs/1905.12322. Showed that bf16, with fp32's range, trains a wide range of models to fp32 quality without loss scaling.
  • Micikevicius et al., FP8 Formats for Deep Learning (2022): https://arxiv.org/abs/2209.05433. Proposed the E4M3 and E5M2 8-bit formats and showed that training and inference hold up in them.
  • Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019): https://arxiv.org/abs/1909.08053. Split each transformer layer's matrix multiplies across the GPUs of one machine, the tensor parallelism of section 5. Annotated companion
  • Sergeev and Del Balso, Horovod: fast and easy distributed deep learning in TensorFlow (2018): https://arxiv.org/abs/1802.05799. Brought the bandwidth-optimal ring all-reduce to deep learning training.

Further reading

on GitHub
   1r"""
   2# The hardware underneath: chips, memory, links and number formats
   3
   4Run: `python -m primer.ml.hardware`
   5
   6New to the notation? `primer.notation` explains every symbol used here from
   7zero. This lesson builds on the cost arithmetic of `primer.ml.inference`.
   8
   9## Level 1: The practitioner's guide
  10
  11**In one sentence.** The hardware under a model is thousands of simple
  12multipliers starved by slow memory, so what you rent or buy is decided by
  13bytes (does the model fit, how fast can its weights be read, how fast can
  14chips talk) and the number format you store those bytes in is the cheapest
  15lever you have.
  16
  17**When you need it.** You need this the day you have to pick a machine: a
  18laptop for local experiments, a cloud GPU for a demo, a multi-GPU server
  19for serving a 70-billion-parameter model, or a cluster for training. You
  20also need it when a model runs far slower than its FLOP count suggests. The
  21tell is a spec sheet you cannot read: teraFLOPS, HBM, NVLink, bf16, FP8, and
  22no idea which number will bite. This lesson's imaginary datacenter GPU
  23does about 10¹⁵ operations per second but reads only 3.35 TB/s, so any work
  24doing fewer than about 299 operations per byte fetched leaves the
  25multipliers idle; generating one token for one user does about 1. You do not
  26need this lesson while you call a hosted API: the provider has done the
  27sizing for you. You need it the moment the bill or the latency makes you
  28consider doing it yourself.
  29
  30**Your options.** From the least hardware to the most, at what each can
  31hold and how its parts talk:
  32
  33| Option | What it is | What fits, roughly | What it costs | Where it lives |
  34|---|---|---|---|---|
  35| A hosted API | Someone else's GPUs behind a per-token price | Any model they offer, at any scale | Money per token, no capacity planning, no control over the machine | The provider |
  36| A laptop or consumer GPU | One chip with a few to a few tens of gigabytes of memory, no fast links | Small models, or larger ones quantized to 4 bits: an 8B model at 4 bits is 4 GB, a 70B is 35 GB (this lesson's memory math) | Cheap and private; slow per token, one user at a time | Your desk |
  37| One datacenter GPU | 80 GB of HBM at 3.35 TB/s (NVIDIA's H100 SXM specification, and this lesson's constants) | A 70B model only at 8 bits (70 GB, 10 GB of cache left) or 4 bits (35 GB, 45 GB left); a 16-bit 70B does not fit | Rental by the hour; the whole card even when one user uses 0.3% of it | A cloud instance or a rack |
  38| One machine, several GPUs on fast links | Chips joined at hundreds of GB/s (NVLink is 900 GB/s on an H100 SXM; this lesson models 500) | A model split across the GPUs, exchanging partial results inside every layer | Several cards' rent; the fast links are what you are paying for | A cloud instance or a rack |
  39| Many machines over a network | Machines joined at tens of GB/s per GPU (this lesson models 50) | Training runs and fleets: each machine holds a copy or a slice, and they talk once per step | The most money and the most engineering; the network becomes the bottleneck | A cluster |
  40
  41**How to choose.** Start from the model's size in bytes and the number
  42format you are willing to run it in.
  43
  44- Compute the weights first: parameters times bytes per weight. If they
  45  fit in one GPU with room for the KV cache, stop there; one chip with no
  46  links is the simplest system you can operate.
  47- If they do not fit, drop the format before adding chips: 8-bit weights
  48  halve the bytes and, on this lesson's numbers, cut the lower bound on
  49  decode time for a 70B model from 41.8 ms to 20.9 ms per token, and 4-bit
  50  to 10.4 ms. Check quality on your own tasks afterwards.
  51- If they still do not fit, add GPUs inside one machine, where the links
  52  are fast enough to split a layer across chips.
  53- Cross to many machines only for training or for a fleet, and design the
  54  split so that the chatty parallelism (tensor parallelism, talking inside
  55  every layer) stays within a machine and only once-per-step traffic
  56  crosses the network.
  57- For training, pick bf16 for the multiplies and keep fp32 master weights:
  58  a gradient of 10⁻⁸ becomes exactly 0 in fp16 but survives in bf16, and a
  59  weight update of 0.001 vanishes in bf16 unless the master copy is fp32
  60  (this lesson's format table). FP8 training is real and works on models up
  61  to 175B parameters with no hyperparameter changes (Micikevicius et al.,
  62  *FP8 Formats for Deep Learning*), but it needs software that handles the
  63  scaling for you.
  64- Whatever you pick, measure what fraction of peak FLOPS you reach. If it
  65  is 80%, you are at least 80% compute-bound; if it is a few percent, you are
  66  moving bytes, and more arithmetic will not help (Horace He, *Making Deep
  67  Learning Go Brrrr*).
  68
  69**What it costs.** Memory is the price of admission and bandwidth is the
  70speed limit. Reading 16 GB of weights once takes 4.78 ms from HBM on this
  71lesson's GPU, 320 ms from the host's memory, 67 times slower: a model
  72"offloaded" to CPU memory runs, but each token waits that much longer.
  73Formats set both bills: halving the bits halves the bytes moved, doubles the
  74operations per byte a tiled kernel achieves (63 becomes 126 in this lesson's
  754096 × 4096 example), and shrinks the multiplier itself (an fp32 multiplier
  76needs 576 cells of silicon, an fp8 one 16), which is why accelerators list
  77roughly double the peak throughput at each halving of the format (the H100
  78lists 1,979 TFLOPS at bf16 and 3,958 at FP8, both with sparsity). What a
  79format costs you in return is range or precision: fp16 tops out at 65,504,
  80fp8 E4M3 at 448, and int4 holds only 15 levels, so small weights vanish
  81without a per-row scale. Links cost time at scale: on this lesson's numbers
  82an all-reduce of a 14 GB gradient across 8 GPUs takes 49 ms inside a machine
  83and 490 ms across machines, against 690 ms of arithmetic per step, so the
  84same run spends 7% of its time talking on fast links and 71% over a network.
  85Power is part of the rent too: an H100 SXM is rated up to 700 W.
  86
  87**What breaks.**
  88
  89- **The model "fits" and then does not.** Weights are the fixed cost; the KV
  90  cache grows with every token of every conversation, and activations and
  91  the serving software take several more gigabytes. Size for weights plus
  92  cache plus headroom, not weights alone.
  93- **A big GPU idles on a small job.** A single user's decode reads every
  94  weight to do two operations with it. Without batching, most of the card
  95  you rent does nothing.
  96- **fp16 training silently zeros gradients**: it loses precision below
  97  6 × 10⁻⁵ and rounds anything under about 3 × 10⁻⁸ to zero (this lesson's
  98  format table). Use loss scaling, or use bf16, which trains to fp32
  99  quality with no hyperparameter changes (Kalamkar et al.).
 100- **bf16 swallows small updates**: 1 + 0.001 rounds back to 1. Keep master
 101  weights and long running sums in fp32.
 102- **fp8 overflows**: 500 in E4M3 is not a number. The format needs per-tensor
 103  scaling that the training or serving library supplies; do not cast by hand.
 104- **Offloading to host memory** makes a model fit at the price of tens of
 105  times slower steps. It is for experiments, not for serving.
 106- **Tensor parallelism across a network** stalls in every layer. Keep it on
 107  the fast links inside a machine.
 108
 109**In the wild.** NVIDIA's H100 specification gives the numbers this
 110lesson rounds (80 GB at 3.35 TB/s, 900 GB/s NVLink). The formats each have a paper: Kalamkar et al. studied bf16 for training,
 111and Micikevicius et al. proposed the two fp8 encodings, E4M3 and E5M2, and
 112earlier the mixed-precision recipe (fp32 master weights, loss scaling) that
 113PyTorch's automatic mixed precision, linked in Further reading, implements.
 114The roofline model this lesson uses to decide
 115memory-bound from compute-bound is Williams, Waterman and Patterson's, and
 116FlashAttention (Dao et al.) is the best-known application of tiling to a
 117model. Megatron-LM (Shoeybi et al.) is the tensor parallelism that lives on
 118fast links, and Horovod (Sergeev and Del Balso) brought the ring all-reduce
 119to deep learning. *How to Scale Your Model*, linked in Further reading,
 120carries the same arithmetic through TPUs and GPUs to full training runs.
 121
 122**Go deeper.** Level 2 builds each number here from nothing: a matrix
 123multiply counted by hand, a memory hierarchy with its six levels, a tiled
 124multiply whose traffic you can watch fall, a 16-bit float encoded bit by bit
 125and every format's range and precision derived from its bit widths, a ring
 126all-reduce simulated on four GPUs, and the serving-fit table computed from
 127the formulas. If you only needed to choose a machine and a format, you are
 128done.
 129
 130## Level 2: How it works, from scratch
 131
 132Every lesson so far has counted operations: so many multiplies per token,
 133so many parameters. This lesson looks at the machine that performs them,
 134because the machine explains things the maths alone never will: why a model
 135that "needs" a tenth of a millisecond of arithmetic takes five milliseconds
 136per token, why training runs are spread across thousands of chips in a
 137particular way, and why everyone is shrinking numbers from 32 bits to 8.
 138
 139Three facts carry the whole lesson:
 140
 1411. **A GPU is thousands of simple arithmetic units** doing the same step on
 142   different numbers. Neural networks are mostly matrix multiplies, which
 143   are exactly that kind of work.
 1442. **Arithmetic is cheap; moving data is expensive.** The chip can multiply
 145   far faster than its memory can feed it, so speed is usually decided by
 146   how many bytes move, not how many operations run.
 1473. **Fewer bits per number helps everywhere at once**: more numbers per
 148   byte moved, more numbers in memory, and smaller, more numerous
 149   multipliers on the chip.
 150
 151Sections 5 and 6 then apply those facts to many GPUs working together, and
 152to the question every deployment starts with: will the model fit?
 153
 154## 1. Why GPUs: thousands of simple cooks
 155
 156**Everyday picture.** A CPU is a few master chefs. Each can cook anything,
 157improvise, and follow a recipe full of "if the sauce splits, do this
 158instead". A GPU is a kitchen of thousands of line cooks who all do the same
 159step at the same moment on different ingredients: "everyone, chop your
 160carrot now." That kitchen is useless for inventing a menu and unbeatable at
 161ten thousand identical salads. A neural network is ten thousand identical
 162salads: nearly all of its work is **matrix multiplication**, the same
 163multiply-and-add done billions of times on different numbers.
 164
 165**Tiny worked example.** Multiply a 2 × 3 matrix by a 3 × 2 matrix:
 166
 167$$
 168A = \begin{pmatrix} 1 & 2 & 3 \\ 4 & 5 & 6 \end{pmatrix}, \quad
 169B = \begin{pmatrix} 7 & 8 \\ 9 & 10 \\ 11 & 12 \end{pmatrix}, \quad
 170AB = \begin{pmatrix} 58 & 64 \\ 139 & 154 \end{pmatrix}
 171$$
 172
 173**Symbols**
 174
 175| Symbol | Meaning here | Shape |
 176|---|---|---|
 177| $A$ | the left matrix: 2 rows of 3 numbers | 2 × 3 |
 178| $B$ | the right matrix: 3 rows of 2 numbers | 3 × 2 |
 179| $AB$ | the **matrix multiply**: the cell in row $i$, column $j$ is row $i$ of $A$ dotted with column $j$ of $B$ | 2 × 2 |
 180| $\begin{pmatrix}\ldots\end{pmatrix}$ | a matrix written out, row by row | |
 181
 182**In words:** "each cell of the answer is one row of A times one column of
 183B, multiplied position by position and added up."
 184
 185**With the numbers:** the top-left cell is 1·7 + 2·9 + 3·11 = 7 + 18 + 33 =
 186**58**; the bottom-right is 4·8 + 5·10 + 6·12 = 32 + 50 + 72 = **154**. Each
 187cell took 3 multiplies and 3 additions (each product added to a running
 188total that starts at 0), so the 4 cells took 12 multiplies and 12 additions.
 189The crucial detail: **no cell needs any other cell's answer.** Four cooks
 190could take one cell each and finish at the same moment.
 191
 192**In Python:**
 193
 194```python
 195A = [[1, 2, 3], [4, 5, 6]]
 196B = [[7, 8], [9, 10], [11, 12]]
 197m, k, n = len(A), len(B), len(B[0])
 198# each cell: row i of A dotted with column j of B
 199[[sum(A[i][p] * B[p][j] for p in range(k)) for j in range(n)] for i in range(m)]  # → [[58, 64], [139, 154]]
 200```
 201
 202That count generalises into the most useful formula in this lesson. A
 203**FLOP** (floating-point operation) is one multiply or one add, and
 204multiplying an m × k matrix by a k × n matrix costs:
 205
 206$$
 207\text{FLOPs} = 2\,m\,n\,k
 208$$
 209
 210**Symbols**
 211
 212| Symbol | Meaning here | In the example |
 213|---|---|---|
 214| $m$ | rows of the left matrix (and of the answer) | 2 |
 215| $k$ | the shared inner size: columns of the left, rows of the right; the length of each dot product | 3 |
 216| $n$ | columns of the right matrix (and of the answer) | 2 |
 217| $m\,n$ | how many cells the answer has | 4 |
 218| $2$ | one multiply plus one add per step of a dot product | |
 219| FLOPs | floating-point operations in total | 24 |
 220
 221**In words:** "every one of the m·n answer cells is a dot product of length
 222k, and every step of a dot product is a multiply and an add."
 223
 224**With the numbers:** 2 × 2 × 2 × 3 = **24** for the example. A 4096 × 4096
 225by 4096 × 4096 multiply, the size of one weight matrix in a mid-sized model,
 226costs 2 × 4096³ ≈ 1.37 × 10¹¹ FLOPs. At the 10¹⁵ FLOPs per second of the
 227imaginary datacenter GPU used throughout `primer.ml.inference`, that is
 228about 0.14 milliseconds, if the chip could be kept busy.
 229
 230**In Python:**
 231
 232```python
 233m, n, k = 2, 2, 3
 234# 2 FLOPs per step, k steps per cell, m·n cells
 2352 * m * n * k  # → 24
 236flops = 2 * 4096 * 4096 * 4096
 237flops  # → 137438953472
 238# milliseconds at 10¹⁵ FLOPs per second
 239round(flops / 1e15 * 1000, 2)  # → 0.14
 240```
 241
 242Hardware usually does "multiply, then add to a running total" as a single
 243instruction, the **fused multiply-add**, which is why FLOPs come in pairs.
 244
 245```mermaid
 246flowchart LR
 247  subgraph CPU["CPU: a few master chefs"]
 248    direction TB
 249    c1["core: big control unit,<br/>big cache, runs any code"]
 250    c2["core"]
 251    c3["core"]
 252  end
 253  subgraph GPU["GPU: thousands of line cooks"]
 254    direction TB
 255    g1["group of cores:<br/>one instruction, many numbers"]
 256    g2["group of cores"]
 257    g3["... about a hundred groups"]
 258    g4["matrix units: a small<br/>tile multiply per instruction"]
 259  end
 260  W["matrix multiply:<br/>m × n independent cells"] --> GPU
 261  BR["branchy code:<br/>if this then that"] --> CPU
 262```
 263
 264**Reading it:** the two boxes spend their silicon differently. The CPU
 265spends it on a few cores that each handle any instruction stream quickly,
 266including code full of decisions. The GPU spends it on arithmetic: groups
 267of cores that all execute the *same* instruction on *different* numbers,
 268plus (on most modern accelerators) **matrix units** that multiply a small
 269tile of numbers in one instruction. A matrix multiply is a perfect fit for
 270the GPU, because its m × n cells are independent and identical in shape.
 271
 272How far does the independence go? Suppose each core takes whole cells and
 273performs one multiply-add per round:
 274
 275$$
 276\text{rounds} = \left\lceil \frac{m\,n}{\text{cores}} \right\rceil \times k
 277$$
 278
 279**Symbols**
 280
 281| Symbol | Meaning here | In the example |
 282|---|---|---|
 283| $m\,n$ | independent cells to compute | 64 × 64 = 4,096 |
 284| cores | identical cores working at once | 1 to 8,192 |
 285| $\lceil x \rceil$ | **ceiling**: round $x$ up to a whole number (half a cell of work still needs a round) | $\lceil 0.5 \rceil = 1$ |
 286| $k$ | multiply-adds per cell, done one after another | 64 |
 287| rounds | how long the whole multiply takes, in rounds | |
 288
 289**In words:** "share the cells out evenly, round up, and each core then
 290spends k rounds on each cell it was given."
 291
 292**With the numbers:** a 64 × 64 by 64 × 64 multiply on 1 core takes
 2934,096 × 64 = **262,144** rounds; on 64 cores, 4,096; on 4,096 cores, **64**.
 294On 8,192 cores it is *still* 64: there are only 4,096 cells, so half the
 295cores have nothing to do.
 296
 297**In Python:**
 298
 299```python
 300import math
 301m = n = k = 64
 302def rounds(cores):
 303    # ⌈m·n / cores⌉ cells per core, then k multiply-adds per cell
 304    return math.ceil(m * n / cores) * k
 305rounds(1), rounds(64), rounds(4096), rounds(8192)  # → (262144, 4096, 64, 64)
 306```
 307
 308![A 64 by 64 multiply speeds up in a straight line from 1 core to 4,096 cores, 262,144 rounds down to 64, then stays flat because there is no more independent work](figures/primer.ml.hardware.parallel_rounds.svg)
 309
 310**Reading it:** both axes are logarithmic. Doubling the cores halves the
 311time, a straight line, until the cores match the number of independent
 312cells (the dashed line at 4,096). Past that point the line goes flat: more
 313cores cannot help a problem that has run out of independent work. A small
 314matrix leaves a big GPU mostly idle.
 315
 316**In code:** `counted_matmul` multiplies with plain loops and counts every multiply and add; `matmul_flops` is the 2·m·n·k formula; `parallel_rounds` is the rounds formula above.
 317
 318**Why it matters in practice.** A GPU is fast only when it is given a lot of
 319independent work at once: large matrices, and many sequences processed
 320together. That is why serving systems batch requests together
 321(`primer.ml.inference`) and why a small model answering one user at a time
 322uses a sliver of the chip. It is also why neural networks look the way they
 323do: architectures that turn into a few big matrix multiplies (the
 324transformer, `primer.ml.transformer`) won partly because they suit this
 325hardware, while step-by-step recurrences (`primer.ml.cnn_rnn`) do not.
 326
 327## 2. The memory hierarchy: near is small, far is big
 328
 329**Everyday picture.** Back in the kitchen. A cook's hands hold one or two
 330things (the **registers**). The cutting board holds a few more (the
 331**on-chip SRAM**, fast memory built into the chip itself). The fridge in
 332the kitchen holds the day's ingredients (**HBM**, "high-bandwidth memory",
 333the GPU's main memory). The storeroom down the hall is the CPU's memory
 334(**host memory**). The warehouse across town is the **disk**. And other
 335restaurants' pantries, reached by courier, are other machines over the
 336**network**. Every step further out holds more and takes longer to reach.
 337The cooks are fast; what slows the kitchen down is fetching.
 338
 339**Tiny worked example.** Round, illustrative numbers for one datacenter GPU
 340and the machine around it (orders of magnitude, not any product's
 341specification):
 342
 343| Level | Holds about | Moves about | Streaming 1 GB takes | What lives there |
 344|---|---|---|---|---|
 345| registers | 20 MB (across the chip) | 100 TB/s | 0.01 ms | the numbers being multiplied this instant |
 346| on-chip SRAM | 50 MB | 20 TB/s | 0.05 ms | tiles of the current multiply (section 3) |
 347| HBM | 80 GB | 3.35 TB/s | 0.30 ms | weights, activations, the KV cache |
 348| host memory | 1 TB | 50 GB/s (over the link to the GPU) | 20 ms | data waiting to be loaded, offloaded state |
 349| local disk | 10 TB | 10 GB/s | 100 ms | datasets, checkpoints |
 350| network | the whole cluster | 50 GB/s per GPU | 20 ms | other GPUs' gradients, remote storage |
 351
 352Registers and SRAM never hold a whole gigabyte; the column shows their
 353*rate*. Notice the jump from HBM to host memory: about **67 times** slower.
 354A GPU that has to reach past its own HBM is a cook walking to the storeroom
 355for every carrot.
 356
 357$$
 358t = \frac{\text{bytes}}{\text{bandwidth}}
 359$$
 360
 361**Symbols**
 362
 363| Symbol | Meaning here | Units |
 364|---|---|---|
 365| $t$ | time to stream the data, ignoring the fixed delay before the first byte arrives (latency) | seconds |
 366| bytes | how much data moves | bytes (1 GB = 10⁹) |
 367| bandwidth | how many bytes per second the level can deliver | bytes per second |
 368
 369**In words:** "the time to move data is its size divided by the speed of
 370the pipe it moves through."
 371
 372**With the numbers:** an 8-billion-parameter model at 2 bytes per
 373parameter is 16 GB. Reading it once from HBM takes 16 × 10⁹ / 3.35 × 10¹² =
 374**4.78 ms**: exactly the lower bound on time per generated token in
 375`primer.ml.inference`, because generating one token reads every weight once.
 376From host memory it would take **320 ms**.
 377
 378**In Python:**
 379
 380```python
 381HBM, host = 3.35e12, 50e9
 382weights = 16e9
 383# t = bytes / bandwidth, in milliseconds
 384round(weights / HBM * 1000, 2)  # → 4.78
 385round(weights / host * 1000)  # → 320
 386# how many times slower the storeroom is than the fridge
 387round(HBM / host)  # → 67
 388```
 389
 390```mermaid
 391flowchart LR
 392  ALU["arithmetic units"] <--> R["registers<br/>~20 MB, ~100 TB/s"]
 393  R <--> S["on-chip SRAM<br/>~50 MB, ~20 TB/s"]
 394  S <--> H["HBM<br/>~80 GB, ~3.35 TB/s"]
 395  H <--> D["host memory<br/>~1 TB, ~50 GB/s"]
 396  D <--> K["local disk<br/>~10 TB, ~10 GB/s"]
 397  H <--> N["network: other machines<br/>~50 GB/s per GPU"]
 398```
 399
 400**Reading it:** start at the arithmetic units on the left and walk outward.
 401Every box to the right is bigger and slower. The arithmetic units can only
 402work on numbers in registers, so every number used must travel the whole
 403way in from wherever it lives. The network hangs off HBM because fast
 404clusters let a GPU send and receive data straight from its own memory
 405without a detour through the CPU.
 406
 407![Capacity grows about fifty-million-fold from registers to the network while bandwidth falls ten-thousand-fold from registers to disk](figures/primer.ml.hardware.hierarchy.svg)
 408
 409**Reading it:** the same six levels, top to bottom, on logarithmic axes.
 410On the left, capacity climbs from megabytes to a petabyte. On the right,
 411bandwidth falls from a hundred terabytes per second to ten gigabytes per
 412second. The one place the order bends is the last two rows: a datacenter
 413network is built to rival a local disk, so fetching from a nearby machine
 414can be as fast as reading your own drive. Both are still around a hundred
 415times slower than HBM.
 416
 417**In code:** `MEMORY_HIERARCHY` lists the six levels with their capacity and bandwidth; `transfer_seconds` is the formula above.
 418
 419**Why it matters in practice.** A model runs at full speed only if
 420everything it touches every step (weights, activations, KV cache) lives in
 421HBM. Spilling to host memory ("offloading") makes a model fit, at the price
 422of each step waiting tens of times longer. And because HBM itself is slow
 423compared with the arithmetic, the fastest code is the code that makes each
 424trip to HBM count, which is the next section.
 425
 426## 3. Arithmetic intensity: why data movement dominates
 427
 428**Everyday picture.** A sandwich shop. If the cook walks to the storeroom
 429for each slice of bread for each sandwich, the cook spends the day walking.
 430If the cook carries a tray of bread and a tray of fillings to the bench and
 431makes a batch of sandwiches from them, every trip feeds many sandwiches.
 432Same sandwiches (FLOPs), far fewer trips (bytes). The ratio of the two is
 433the **arithmetic intensity**: operations done per byte fetched.
 434
 435Our imaginary GPU does 10¹⁵ FLOPs per second but reads only 3.35 × 10¹²
 436bytes per second from HBM, so it breaks even at about **299 FLOPs per
 437byte** (the ridge point of the roofline in `primer.ml.inference`). Any work
 438doing fewer operations than that for each byte it fetches leaves the
 439arithmetic units waiting on memory.
 440
 441**Tiny worked example.** Multiply two 4 × 4 matrices: 2 × 4³ = 128 FLOPs.
 442Count the numbers fetched from slow memory (HBM) into fast memory (SRAM):
 443
 444| Strategy | Numbers read | Numbers written | FLOPs per number moved |
 445|---|---|---|---|
 446| no reuse: each cell fetches its own row of A and column of B | 16 × (4 + 4) = **128** | 16 | 128 / 144 = 0.89 |
 447| 2 × 2 tiles: load a tile of A and a tile of B, use each number twice | **64** | 16 | 128 / 80 = 1.6 |
 448| one 4 × 4 tile: load everything once | **32** | 16 | 128 / 48 = 2.7 |
 449
 450The arithmetic is identical in all three rows. Only the traffic changes.
 451This trick is **tiling**: bring a small block of each matrix into fast
 452memory and do every multiplication that block takes part in before
 453throwing it away.
 454
 455$$
 456\text{reads} = \frac{2\,n^3}{T}
 457\qquad
 458I = \frac{2\,n^3}{b\left(\dfrac{2\,n^3}{T} + n^2\right)} \approx \frac{T}{b}
 459$$
 460
 461**Symbols**
 462
 463| Symbol | Meaning here | In the example |
 464|---|---|---|
 465| $n$ | both matrices are $n \times n$ | 4, then 4,096 |
 466| $T$ | tile width: fast memory works on $T \times T$ blocks | 1, 2, 4, then 128 |
 467| $2\,n^3$ | the FLOPs of the whole multiply (section 1, with $m = n = k$) | 128 |
 468| reads | numbers fetched from slow memory | 128, 64, 32 |
 469| $n^2$ | numbers written back: each answer cell once | 16 |
 470| $b$ | bytes per number: 2 at 16-bit, 1 at 8-bit | 2 |
 471| $I$ | arithmetic intensity: FLOPs per byte moved | FLOPs/byte |
 472| $\approx$ | "roughly", once $n$ is much bigger than $T$ and the writes are negligible | |
 473
 474**In words:** "every number fetched is used T times, so the traffic falls in
 475proportion to the tile width, and the intensity rises in proportion to it."
 476
 477**With the numbers:** for n = 4, the reads are 2 × 64 / T = 128, 64 and 32
 478for tiles of 1, 2 and 4, as the table says. For two 4096 × 4096 matrices in
 47916-bit with 128-wide tiles, I ≈ 128 / 2 = 64, and **63.0** once the
 480writes are counted. In 8-bit the same tiles give **126**: halving the bytes
 481per number doubles the intensity.
 482
 483**In Python:**
 484
 485```python
 486n = 4
 487# reads = 2n³ / T, for tiles of 1, 2 and 4
 488[2 * n**3 // T for T in (1, 2, 4)]  # → [128, 64, 32]
 489def intensity(n, T, b):
 490    flops = 2 * n**3
 491    # bytes moved: every read, plus one write per answer cell
 492    moved = b * (2 * n**3 / T + n**2)
 493    return flops / moved
 494round(intensity(4096, 128, 2), 1)  # → 63.0
 495round(intensity(4096, 128, 1), 1)  # → 126.0
 496```
 497
 498```mermaid
 499flowchart LR
 500  subgraph HBM["HBM: big, slow"]
 501    A["A, in T × T tiles"]
 502    B["B, in T × T tiles"]
 503    C["C, the answer"]
 504  end
 505  subgraph SRAM["on-chip SRAM: small, fast"]
 506    a["one tile of A"]
 507    b["one tile of B"]
 508    acc["running total for<br/>one T × T tile of C"]
 509  end
 510  A -- "load" --> a
 511  B -- "load" --> b
 512  a --> mm["multiply-add:<br/>T³ steps, no traffic"]
 513  b --> mm
 514  mm --> acc
 515  mm -. "next pair of tiles along k" .-> A
 516  acc -- "write once, at the end" --> C
 517```
 518
 519**Reading it:** the left box is slow memory, the right box fast memory.
 520For one tile of the answer, the kernel repeatedly loads one tile of A and
 521one tile of B, does all T³ multiply-adds between them without touching slow
 522memory, and adds the results into a running total that stays on chip. Only
 523when the whole row of tiles has been consumed does it write the finished
 524answer tile back, once. Fast memory never holds more than three tiles.
 525
 526![Measured reads fall from 65,536 to 2,048 as the tile grows from 1 to 32, on the 2n-cubed-over-T line, and intensity at n = 4096 rises with the tile, reaching the 299 break-even only near T = 650 in 16-bit or T = 310 in 8-bit](figures/primer.ml.hardware.tiling.svg)
 527
 528**Reading it:** on the left, dots are reads counted by actually running the
 529tiled multiply on 32 × 32 matrices, and the line is the formula 2n³/T; they
 530agree exactly, and each doubling of the tile halves the traffic. On the
 531right, the intensity of a 4096 × 4096 multiply climbs with the tile width,
 532and the dashed line is the chip's break-even point of 299. In 16-bit the
 533tiles need to be about 650 wide to cross it. Three 650 × 650 tiles in 16-bit
 534take about 2.5 MB, more fast memory than one group of cores has to itself,
 535which is why real kernels tile at several levels at once (tiles in
 536registers inside tiles in SRAM) and why lower-precision numbers, which
 537shift the whole curve up, are so attractive.
 538
 539The same idea runs through the rest of this primer:
 540
 541- **FlashAttention** (`primer.ml.attention`) tiles attention: blocks of
 542  queries, keys and values are loaded into SRAM, and the n × n score matrix
 543  is never written to HBM at all. Same answer, a fraction of the traffic.
 544- **Kernel fusion**: adding a bias or applying an activation does about one
 545  FLOP per number it reads, hopelessly below 299, so these steps are done
 546  inside the matrix-multiply kernel while the tile is still on chip.
 547- **Decode** (`primer.ml.inference`): generating one token for one user
 548  reads every weight to do just 2 FLOPs with it, an intensity of about 1.
 549  Batching users together is tiling across requests: one read of a weight
 550  serves every sequence in the batch.
 551
 552**In code:** `tiled_matmul` runs the tiled multiply and returns a `Traffic` count of reads, writes, FLOPs and peak fast-memory use; `matmul_reads` and `matmul_intensity` are the formulas; `primer.ml.inference.ridge_point` is the break-even.
 553
 554**Why it matters in practice.** Before asking how many FLOPs a piece of
 555work needs, ask how many bytes it moves and how often each byte is reused.
 556Most large speed-ups in modern AI systems (FlashAttention, fused kernels,
 557batching, quantization) change the bytes, not the FLOPs.
 558
 559## 4. Number formats: how many bits each number gets
 560
 561**Everyday picture.** Scientific notation on a form with a fixed number of
 562boxes: 6.02 × 10²³. One box holds the sign. A few boxes hold the power of
 563ten, which sets how big or small the number can be: its **range**. The
 564rest hold the digits, which set how finely it is measured: its
 565**precision**. With a fixed number of boxes, moving a box from the digits
 566to the power buys range and costs precision. Computers do the same with
 567bits and powers of two, and call the parts the **sign**, the **exponent**
 568and the **mantissa** (the stored digits).
 569
 570**Tiny worked example.** Store −6.5 in **bf16** ("brain float 16": 1 sign
 571bit, 8 exponent bits, 7 mantissa bits).
 572
 5731. **Sign:** negative, so the sign bit is 1.
 5742. **Power of two:** the largest power of two not above 6.5 is 4 = 2², so
 575   6.5 = 1.625 × 2².
 5763. **Exponent:** stored with a **bias** of 127 added, so that negative
 577   powers need no sign of their own: 2 + 127 = 129 = 10000001 in binary.
 5784. **Mantissa:** the leading 1 of 1.625 is always there, so it is not
 579   stored. The fraction 0.625 = ½ + ⅛ is 0.101 in binary, padded to seven
 580   bits: 1010000.
 5815. **The 16 bits:** 1 10000001 1010000.
 582
 583Most numbers are not so lucky. 0.1 has no finite binary expansion, so it is
 584rounded to the nearest value each format can hold: 0.10000000149 in fp32,
 5850.1000977 in bf16, 0.0999756 in fp16 and 0.1015625 in 8-bit E4M3.
 586
 587$$
 588x = (-1)^{s} \times 2^{\,e - \text{bias}} \times \left(1 + \frac{f}{2^{M}}\right),
 589\qquad \text{bias} = 2^{E-1} - 1
 590$$
 591
 592**Symbols**
 593
 594| Symbol | Meaning here | In the example |
 595|---|---|---|
 596| $s$ | the sign bit: 0 positive, 1 negative | 1 |
 597| $(-1)^s$ | −1 multiplied by itself $s$ times: +1 when $s = 0$, −1 when $s = 1$ | −1 |
 598| $E$ | how many exponent bits the format has | 8 |
 599| $e$ | the stored exponent, read as an ordinary whole number | 129 |
 600| bias | the offset subtracted from $e$, so stored values 1…254 stand for powers −126…127 | 127 |
 601| $M$ | how many mantissa bits the format has | 7 |
 602| $f$ | the stored mantissa, read as a whole number from 0 to $2^M - 1$ | 1010000 = 80 |
 603| $1 + f/2^M$ | the significand: the hidden leading 1 plus the stored fraction, between 1 and 2 | 1 + 80/128 = 1.625 |
 604| $x$ | the number the bits stand for | −6.5 |
 605
 606**In words:** "the sign says plus or minus, the exponent says which power of
 607two to scale by, and the mantissa says how far between that power and the
 608next one the number sits."
 609
 610**With the numbers:** (−1)¹ × 2^(129 − 127) × (1 + 80/128) = −1 × 4 × 1.625 =
 611**−6.5**.
 612
 613**In Python:**
 614
 615```python
 616x = -6.5
 617M, bias = 7, 127
 618# s: 1 for a negative number
 619s = 1 if x < 0 else 0
 620# e: the power of two below |x| is 2², stored with the bias added
 621e = 2 + bias
 622e, format(e, "08b")  # → (129, '10000001')
 623# f: the fraction after the hidden 1 of |x| / 2², as a 7-bit whole number
 624f = round((abs(x) / 2**2 - 1) * 2**M)
 625f, format(f, "07b")  # → (80, '1010000')
 626# decode: (-1)^s × 2^(e - bias) × (1 + f / 2^M)
 627(-1)**s * 2**(e - bias) * (1 + f / 2**M)  # → -6.5
 628```
 629
 630Two corners of the formula matter in practice. When the stored exponent is
 6310, the hidden 1 is dropped and the number is a **subnormal**: it lets values
 632fade gradually towards zero instead of dropping off a cliff, at the cost of
 633fewer significant bits. And IEEE-style formats reserve the all-ones
 634exponent for **infinity** and **NaN** ("not a number"), which is where
 635overflowing values go.
 636
 637```mermaid
 638flowchart LR
 639  X["x = −6.5"] --> S["sign: negative<br/>s = 1"]
 640  X --> P["largest power of two<br/>not above 6.5: 2² = 4"]
 641  P --> E["exponent: 2 + bias 127<br/>e = 129 = 10000001"]
 642  P --> F["6.5 / 4 = 1.625<br/>drop the leading 1: .625"]
 643  F --> R["round .625 to 7 bits<br/>f = 1010000"]
 644  S --> B["1 | 10000001 | 1010000"]
 645  E --> B
 646  R --> B
 647```
 648
 649**Reading it:** a number enters on the left and splits three ways. The sign
 650is read off directly. The exponent comes from finding which pair of powers
 651of two the number sits between. The mantissa is where the number sits
 652within that pair, rounded to however many bits the format allows: this
 653rounding box is the only place information is lost, and it is where every
 654format differs.
 655
 656![Bit layouts drawn to scale: fp32 has 1 sign, 8 exponent and 23 mantissa bits; bf16 keeps the 8 exponent bits and cuts the mantissa to 7; fp16 has 5 and 10; the two fp8 formats have 5 and 2, or 4 and 3; int8 and int4 are plain integers](figures/primer.ml.hardware.format_layouts.svg)
 657
 658**Reading it:** each row is one format drawn to scale, one cell per bit,
 659with the sign in grey, the exponent in orange and the mantissa in blue.
 660Compare bf16 and fp16: the same 16 bits, split differently. bf16 is
 661literally the top half of fp32, keeping all 8 exponent bits (fp32's whole
 662range) and giving up precision; fp16 keeps more precision and gives up
 663range. The integer formats at the bottom have no exponent at all: every
 664value is a whole number, turned into a weight by one shared scale
 665(`primer.ml.inference` builds int8 and int4 quantization from scratch).
 666
 667| Format | Bits (sign, exponent, mantissa) | Largest | Smallest at full precision | Gap just above 1 | Typical use |
 668|---|---|---|---|---|---|
 669| fp32 | 1, 8, 23 | 3.4 × 10³⁸ | 1.2 × 10⁻³⁸ | 1.2 × 10⁻⁷ | master weights, optimizer state, running sums |
 670| bf16 | 1, 8, 7 | 3.4 × 10³⁸ | 1.2 × 10⁻³⁸ | 1/128 ≈ 0.0078 | training and inference matrix multiplies |
 671| fp16 | 1, 5, 10 | 65,504 | 6.1 × 10⁻⁵ | 1/1024 ≈ 0.00098 | inference; training with loss scaling |
 672| fp8 E5M2 | 1, 5, 2 | 57,344 | 6.1 × 10⁻⁵ | 0.25 | gradients in 8-bit training |
 673| fp8 E4M3 | 1, 4, 3 | 448 | 0.0156 | 0.125 | weights and activations in 8-bit |
 674| int8 | 8-bit integer | 127 × scale | evenly spaced | none: a fixed step | quantized weights |
 675| int4 | 4-bit integer | 7 × scale | evenly spaced | none: a fixed step | quantized weights |
 676
 677E4M3 bends the IEEE rules: it has no infinity, and spends that exponent on
 678ordinary numbers instead, which is how 8 bits reach 448 rather than 240.
 679
 680![Relative spacing between neighbouring values: each float format is a flat band across its range (fp32 near 1e-7, bf16 near 1e-2, fp8 near 0.1), rising at its small end and stopping at its largest value, while int8 and int4 spacing rises steadily as numbers shrink](figures/primer.ml.hardware.precision.svg)
 681
 682**Reading it:** the x-axis is the size of the number being stored; the
 683y-axis is the gap to the next storable value, as a fraction of the number
 684(lower is more precise); both are logarithmic. Each float format is a flat
 685band: its relative precision is the same for tiny and huge numbers, which
 686is the whole point of an exponent. The band's width is the range: it stops
 687on the right at the largest value (65,504 for fp16, 448 for E4M3) and rises
 688on the left where subnormals run out of bits. fp32 and bf16 are flat across
 689the whole plot and far beyond it. The integer formats, here with a scale
 690that maps 1.0 to the top code, are straight rising lines: their step is the
 691same size everywhere, so small numbers get coarse, which is why quantized
 692models use one scale per row or block of weights.
 693
 694What the picture means for training (`primer.ml.pretraining` covers mixed
 695precision in full):
 696
 697- **Range failures.** A gradient of 10⁻⁸ becomes exactly 0 in fp16 (below
 698  its smallest subnormal) but survives in bf16 as 1.0012 × 10⁻⁸. fp16
 699  training therefore multiplies the loss by a large constant (**loss
 700  scaling**) to lift gradients into range; bf16 training does not need to.
 701- **Precision failures.** In bf16, 1 + 0.001 rounds back to exactly 1: a
 702  small weight update simply vanishes. So training keeps a **master copy**
 703  of the weights in fp32, adds its sums up in fp32, and uses 16-bit (or
 704  8-bit) only for the big multiplies. That split is **mixed precision**.
 705
 706### Why smaller formats multiply throughput
 707
 708Fewer bits per number pays three times:
 709
 7101. **Bytes.** Half the bytes per number moves twice the numbers per second
 711   through every level of section 2, and fits twice the parameters in HBM.
 712   Memory-bound work speeds up directly: the lower bound on time per token
 713   for a 70-billion-parameter model (`primer.ml.inference`) is 41.8 ms at
 714   16-bit, 20.9 ms at 8-bit and 10.4 ms at 4-bit.
 7152. **Intensity.** The same tile does twice the FLOPs per byte (section 3:
 716   63 becomes 126), pushing more work past the break-even point.
 7173. **Silicon.** A multiplier is the expensive part of the chip, and its
 718   size grows with the square of the number of significand bits.
 719
 720**Everyday picture.** Long multiplication by hand: multiplying two 3-digit
 721numbers means writing a 3 × 3 grid of single-digit products; two 6-digit
 722numbers need a 6 × 6 grid, four times the work. A chip's multiplier is that
 723grid built in wires.
 724
 725$$
 726\text{cells} = p^{2}, \qquad p = M + 1
 727$$
 728
 729**Symbols**
 730
 731| Symbol | Meaning here | In the example |
 732|---|---|---|
 733| $M$ | stored mantissa bits | 23, 10, 7, 3 |
 734| $p$ | significand bits: the stored mantissa plus the hidden leading 1 | 24, 11, 8, 4 |
 735| cells | one-bit products in a schoolbook (array) multiplier: one per pair of bits | |
 736
 737**In words:** "multiplying two p-bit significands needs one small cell for
 738every pair of bits, so p times p cells; the exponents only need adding,
 739which is cheap."
 740
 741**With the numbers:** fp32 needs 24² = **576** cells, fp16 11² = 121, bf16
 7428² = 64, and fp8 E4M3 4² = **16**: one fp32 multiplier's worth of silicon
 743holds many 8-bit ones. Accelerators typically list roughly double the peak
 744operations per second each time the format halves.
 745
 746**In Python:**
 747
 748```python
 749# stored mantissa bits for fp32, fp16, bf16 and fp8 E4M3
 750mantissa_bits = [23, 10, 7, 3]
 751# p = M + 1, and cells = p²
 752[(M + 1) ** 2 for M in mantissa_bits]  # → [576, 121, 64, 16]
 753```
 754
 755**In code:** `FloatFormat` describes a format by its exponent and mantissa bits, with `FloatFormat.max_value`, `FloatFormat.min_normal`, `FloatFormat.min_subnormal` and `FloatFormat.epsilon` derived from them; `FP32`, `BF16`, `FP16`, `FP8_E5M2` and `FP8_E4M3` are the five formats; `encode` rounds a number into its three fields with plain arithmetic, `decode` turns fields back into a number, `round_to` does both, and `bit_string` prints the bits; `multiplier_cells` is the p² formula. The tests check `round_to` against NumPy's own float16 and float32.
 756
 757**Why it matters in practice.** Picking a number format is picking a point
 758on the range-versus-precision curve for each kind of number in a model.
 759Weights and activations tolerate coarse formats; gradients need range; the
 760running sums of long dot products need precision. Modern training and
 761serving use a different format for each, and the savings are among the
 762largest in the field.
 763
 764## 5. Many GPUs: the cost of talking
 765
 766**Everyday picture.** A group project. Four people each work through a
 767quarter of the exercises, and then must agree on one combined answer
 768sheet. Sitting at the same table they can compare notes in seconds; living
 769in different cities they must post letters. The more often a team must
 770compare notes, the more it matters who sits at the same table.
 771
 772Large models are trained on many GPUs because no single one has the memory
 773or the speed. The simplest split is **data parallelism**: every GPU holds a
 774copy of the model and works on different examples, and after each step
 775their gradients must be added up, so every copy takes the same step. That
 776"add up, and give everyone the total" operation is an **all-reduce**.
 777
 778**Tiny worked example: a ring all-reduce.** Four GPUs each hold a gradient
 779of 8 numbers. Each splits its gradient into 4 chunks of 2 numbers and they
 780sit in a ring, each passing to its right-hand neighbour.
 781
 7821. **Reduce-scatter, 3 steps:** each GPU sends one chunk to its neighbour,
 783   which adds it to its own copy of that chunk. After 3 steps, each GPU
 784   holds one chunk that contains the sum from all four.
 7852. **All-gather, 3 more steps:** the finished chunks travel round the ring,
 786   overwriting the stale copies.
 787
 788Each GPU sent 6 chunks of 2 numbers, **12 numbers**, which is
 7892 × 3/4 × 8. The remarkable part: with 400 GPUs instead of 4, each would
 790still send just under twice its gradient. The work per GPU barely grows.
 791
 792```mermaid
 793flowchart LR
 794  G0["GPU 0<br/>chunks a b c d"] -- "one chunk per step" --> G1["GPU 1<br/>chunks a b c d"]
 795  G1 -- "one chunk per step" --> G2["GPU 2<br/>chunks a b c d"]
 796  G2 -- "one chunk per step" --> G3["GPU 3<br/>chunks a b c d"]
 797  G3 -- "one chunk per step" --> G0
 798```
 799
 800**Reading it:** four GPUs in a ring, each only ever talking to its
 801right-hand neighbour. Every link is busy at every step, each carrying a
 802different chunk, so no link sits idle and no single GPU becomes a
 803bottleneck. After 2 × (4 − 1) = 6 steps, every GPU has every chunk summed
 804over all four.
 805
 806$$
 807t_{\text{all-reduce}} = \frac{2\,(N-1)}{N} \cdot \frac{S}{B}
 808\qquad
 809t_{\text{compute}} = \frac{6\,P\,D}{F}
 810$$
 811
 812**Symbols**
 813
 814| Symbol | Meaning here | In the example |
 815|---|---|---|
 816| $N$ | GPUs taking part | 8 |
 817| $S$ | bytes of gradient each GPU holds | 7 × 10⁹ parameters × 2 bytes = 14 GB |
 818| $B$ | each GPU's link bandwidth | 500 GB/s in one machine, 50 GB/s between machines (illustrative) |
 819| $\frac{2(N-1)}{N}$ | the fraction of its gradient each GPU sends in a ring all-reduce: just under 2 | 1.75 |
 820| $P$ | parameters in the model | 7 × 10⁹ |
 821| $D$ | tokens each GPU processes per step | 16,384 |
 822| $6$ | FLOPs per parameter per token in training: 2 forward, 4 backward | |
 823| $F$ | the GPU's arithmetic speed | 10¹⁵ FLOPs per second |
 824
 825**In words:** "talking takes just under twice the gradient's size divided by
 826the link speed, however many GPUs there are; computing takes six operations
 827per parameter per token, divided by the chip's speed."
 828
 829**With the numbers:** inside one machine, 1.75 × 14 × 10⁹ / 500 × 10⁹ =
 830**49 ms**. Across machines, **490 ms**. The arithmetic for the step is
 8316 × 7 × 10⁹ × 16,384 / 10¹⁵ = **0.69 s**. Inside a machine, talking costs 7%
 832of the computing time; across machines, 71%.
 833
 834**In Python:**
 835
 836```python
 837N, S = 8, 14e9
 838in_machine, between_machines = 500e9, 50e9
 839# 2(N - 1)/N × S / B, in milliseconds
 840round(2 * (N - 1) / N * S / in_machine * 1000)  # → 49
 841round(2 * (N - 1) / N * S / between_machines * 1000)  # → 490
 842P, D, F = 7e9, 16_384, 1e15
 843# 6·P·D / F, in seconds
 844round(6 * P * D / F, 2)  # → 0.69
 845```
 846
 847![All-reduce time levels off as GPUs are added: about 56 ms over fast in-machine links and about 560 ms over the network, against 690 ms of arithmetic per step](figures/primer.ml.hardware.communication.svg)
 848
 849**Reading it:** the x-axis is the number of GPUs (log scale); the y-axis
 850is seconds per training step. The flat grey line is the arithmetic every
 851GPU does per step. The two rising curves are the all-reduce over each kind
 852of link, and both flatten out almost at once: that is the ring's
 8532(N − 1)/N approaching 2. The gap between the curves is the whole story:
 854the same gradient costs ten times more over the network than over in-machine
 855links, bringing communication close to the cost of the arithmetic itself.
 856Real systems hide part of it by sending the gradients of later layers while
 857earlier layers are still computing theirs.
 858
 859```mermaid
 860flowchart TB
 861  subgraph M1["machine 1: fast links, ~500 GB/s"]
 862    a1["GPU"] <--> a2["GPU"] <--> a3["GPU"] <--> a4["GPU"]
 863  end
 864  subgraph M2["machine 2: fast links, ~500 GB/s"]
 865    b1["GPU"] <--> b2["GPU"] <--> b3["GPU"] <--> b4["GPU"]
 866  end
 867  M1 <-- "network, ~50 GB/s per GPU:<br/>data parallelism, once per step" --> M2
 868  TP["tensor parallelism:<br/>talks inside every layer"] -.-> M1
 869  TP -.-> M2
 870```
 871
 872**Reading it:** two machines, each with a few GPUs joined by fast links, and
 873a slower network between them. The chattiest kind of splitting, **tensor
 874parallelism** (cutting each matrix multiply across GPUs, which must swap
 875partial results inside every layer), is kept within one machine's fast
 876links. Kinds that talk rarely, like data parallelism (one all-reduce per
 877step) and **pipeline parallelism** (passing activations only between
 878neighbouring stages of layers), span the slower network. The layout of the
 879wires decides the layout of the work. `primer.ml.pretraining` builds these
 880strategies in full.
 881
 882**In code:** `ring_all_reduce` simulates the ring step by step and counts the numbers each GPU sends; `all_reduce_seconds` and `training_step_seconds` are the two formulas; `IN_MACHINE_LINK` and `BETWEEN_MACHINES_LINK` are the illustrative link speeds.
 883
 884**Why it matters in practice.** At scale, a cluster's network matters as
 885much as its chips. Parallelism strategies are chosen by matching how often
 886each kind talks to how fast each link is, and a training run that ignores
 887the map spends its budget waiting.
 888
 889## 6. Will it fit? Memory math for serving
 890
 891**Everyday picture.** A bookshelf of fixed width. The encyclopedia (the
 892model's weights) must go on it, whole. Whatever space is left holds one
 893notebook per customer being served (their KV cache), and a notebook grows
 894with every word of the conversation. Thinner volumes, printed in a smaller
 895number format, leave room for more notebooks.
 896
 897**Tiny worked example.** A model shaped like Llama 3 70B (80 layers, 8
 898key/value heads of 128 dimensions each) on one 80 GB GPU. Its KV cache
 899costs 2 × 80 × 8 × 128 × 2 = 327,680 bytes per token at 16-bit
 900(`primer.ml.inference` derives this).
 901
 902| Weights | Weight memory | Left for the cache | Tokens of 16-bit cache | Tokens of 8-bit cache |
 903|---|---|---|---|---|
 904| 16-bit | 140 GB | none: does not fit | 0 | 0 |
 905| 8-bit (fp8) | 70 GB | 10 GB | 30,517 | 61,035 |
 906| 4-bit | 35 GB | 45 GB | 137,329 | 274,658 |
 907
 908This ignores activations and the serving software's own working memory,
 909which take several more gigabytes in practice.
 910
 911$$
 912P \, b_w + C \times 2\,L\,H_{kv}\,d_h\,b_{kv} \le M_{\text{GPU}}
 913$$
 914
 915**Symbols**
 916
 917| Symbol | Meaning here | In the example |
 918|---|---|---|
 919| $P$ | parameters | 7 × 10¹⁰ |
 920| $b_w$ | bytes per weight: bits ÷ 8 | 2, 1 or 0.5 |
 921| $C$ | tokens held in the KV cache, summed over every request being served | the unknown |
 922| $2$ | one key and one value per token | |
 923| $L$ | layers, each with its own cache | 80 |
 924| $H_{kv}$ | key/value heads | 8 |
 925| $d_h$ | dimensions per head | 128 |
 926| $b_{kv}$ | bytes per cached number | 2 or 1 |
 927| $M_{\text{GPU}}$ | the GPU's memory | 80 × 10⁹ bytes |
 928| $\le$ | "must be at most" | |
 929
 930**In words:** "the weights, plus every cached token's keys and values
 931across every layer, must fit in the GPU's memory."
 932
 933**With the numbers:** with fp8 weights, 7 × 10¹⁰ × 1 = 70 GB, leaving 10 GB.
 93410 × 10⁹ / 327,680 = **30,517** tokens of 16-bit cache, or **61,035** with
 935the cache in fp8 too: one long conversation, or a dozen short ones.
 936
 937**In Python:**
 938
 939```python
 940P, M_gpu = 70e9, 80e9
 941L, H_kv, d_h = 80, 8, 128
 942def kv_per_token(b_kv):
 943    # a key and a value, per layer, per KV head
 944    return 2 * L * H_kv * d_h * b_kv
 945kv_per_token(2)  # → 327680
 946# weight GB at 16, 8 and 4 bits
 947[P * bits / 8 / 1e9 for bits in (16, 8, 4)]  # → [140.0, 70.0, 35.0]
 948# tokens of cache beside fp8 weights, with a 16-bit cache and then an fp8 one
 949int((M_gpu - P * 1) // kv_per_token(2)), int((M_gpu - P * 1) // kv_per_token(1))  # → (30517, 61035)
 950```
 951
 952```mermaid
 953flowchart LR
 954  P["parameters × bytes per weight"] --> W["weights"]
 955  T["tokens × bytes per token"] --> K["KV cache"]
 956  W --> Q{"weights + cache<br/>fit in GPU memory?"}
 957  K --> Q
 958  Q -- yes --> Y["serve: spare room means<br/>more users or longer contexts"]
 959  Q -- no --> N["smaller formats, fewer KV heads,<br/>shorter contexts, or more GPUs"]
 960```
 961
 962**Reading it:** two budgets feed one question. The weights are a fixed cost,
 963set by the parameter count and the number format. The cache grows with
 964every token of every active conversation. If they fit together, whatever is
 965left over becomes capacity: more concurrent users, longer contexts. If
 966they don't, every lever on the right shrinks one of the two inputs, or
 967splits the model across GPUs joined by the fast links of section 5.
 968
 969![Stacked bars against an 80 GB line: 16-bit weights alone reach 140 GB and overflow; fp8 weights take 70 GB and leave 10 GB, about 30,000 tokens of cache; 4-bit weights take 35 GB and leave 45 GB, about 137,000 tokens](figures/primer.ml.hardware.serving_fit.svg)
 970
 971**Reading it:** each bar is one choice of weight format for the same
 97270-billion-parameter model. The dark part is the weights; the light part
 973is the room left for the KV cache, labelled with how many 16-bit tokens fit
 974in it. The dashed line is the GPU's 80 GB. At 16-bit the weights alone
 975break through the line. Each halving of the weight format frees memory that
 976turns directly into tokens of cache, which is to say into users.
 977
 978**In code:** `max_cache_tokens` subtracts the weights and divides what is left by the per-token cache, using `primer.ml.inference.weight_bytes` and `primer.ml.inference.kv_cache_bytes_per_token`; `LLAMA3_70B_SHAPE` holds the example model's shape.
 979
 980**Why it matters in practice.** This arithmetic comes first in every
 981deployment: the number format sets whether a model fits, how many users
 982share a GPU, and (because decode reads every weight per token) how fast
 983each of them sees words appear. `primer.ml.inference` carries it on into
 984batching, speculative decoding and prompt caching.
 985
 986## In 20 seconds
 987
 988- **GPUs** are thousands of simple cores doing the same step on different
 989  numbers. Matrix multiplies (2·m·n·k FLOPs, every cell independent) are
 990  the perfect workload, as long as there is enough independent work.
 991- **Memory hierarchy:** registers, on-chip SRAM, HBM, host memory, disk,
 992  network. Each step out is bigger and slower; the GPU runs at full speed
 993  only on data in HBM or closer.
 994- **Arithmetic intensity:** the chip does hundreds of operations in the
 995  time it fetches one byte, so speed is set by bytes moved. Tiling reuses
 996  each fetched number T times and cuts traffic T-fold; FlashAttention,
 997  kernel fusion and batching are the same idea.
 998- **Number formats:** sign, exponent (range) and mantissa (precision).
 999  bf16 keeps fp32's range with less precision; fp16 the reverse; fp8 and
1000  int8/int4 trade more. Fewer bits means fewer bytes, higher intensity and
1001  smaller multipliers, so throughput roughly doubles per halving.
1002- **Many GPUs:** a ring all-reduce costs about 2 × gradient size ÷ link
1003  speed, whatever the GPU count. Chatty parallelism stays on fast
1004  in-machine links; the rest crosses the network.
1005- **Serving:** weights plus KV cache must fit; the number format decides
1006  both what fits and how fast each token comes.
1007
1008## Self-test questions
1009
1010**Why are GPUs, rather than CPUs, used to train and run neural networks?**
1011Nearly all of a network's work is matrix multiplication, where every output
1012cell is an independent dot product. A GPU spends its silicon on thousands
1013of simple arithmetic units (plus matrix units) that apply the same
1014instruction to different numbers, so it can compute thousands of cells at
1015once. A CPU spends its silicon on a few flexible cores that are better at
1016branchy, sequential code.
1017
1018**How many FLOPs does multiplying a 1,000 × 2,000 matrix by a 2,000 × 500 matrix take, and why that formula?**
10192 × 1,000 × 500 × 2,000 = 2 × 10⁹. There are m·n = 500,000 output cells,
1020each a dot product of length k = 2,000, and each step of a dot product is
1021one multiply and one add.
1022
1023**The chip can do 10¹⁵ FLOPs per second but a model runs far slower. What is usually the bottleneck, and how do you tell?**
1024Memory bandwidth. Compare the work's arithmetic intensity (FLOPs per byte
1025moved from HBM) with the chip's break-even ratio, peak FLOPs divided by
1026bandwidth (about 299 here). Below it, the arithmetic units wait on memory;
1027generating one token for one user has an intensity near 1, so it is
1028deeply memory-bound.
1029
1030**How does tiling a matrix multiply reduce memory traffic, and what limits it?**
1031Each block of numbers is loaded into fast on-chip memory once and used for
1032every multiplication it takes part in, T times, instead of being fetched
1033again for each one. Reads fall from 2n³ to 2n³/T. The limit is fast-memory
1034size: three T × T tiles must fit at once, so real kernels tile at several
1035levels of the hierarchy.
1036
1037**What is the difference between bf16 and fp16, and why do many training runs prefer bf16?**
1038Both have 16 bits. bf16 has fp32's 8 exponent bits and 7 mantissa bits:
1039the same range as fp32, less precision. fp16 has 5 exponent bits and 10
1040mantissa bits: more precision, but a range that tops out at 65,504 and
1041rounds gradients smaller than about 3 × 10⁻⁸ to zero. bf16 avoids the need for
1042loss scaling; its coarse precision is handled by keeping fp32 master
1043weights and fp32 sums.
1044
1045**Why does halving the bits per number roughly double throughput?**
1046Three reasons: half the bytes to move, so memory-bound work runs twice as
1047fast and twice the parameters fit; twice the FLOPs per byte for the same
1048tile, so more work clears the break-even point; and multipliers whose area
1049grows with the square of the significand bits, so many more small
1050multipliers fit in the same silicon.
1051
1052**Why is tensor parallelism usually kept inside one machine while data parallelism spans many?**
1053Tensor parallelism splits every matrix multiply, so GPUs must exchange
1054partial results inside every layer, many times per step: it needs the fast
1055in-machine links. Data parallelism communicates once per step (an
1056all-reduce of the gradients), which a ring spreads so each GPU sends only
1057about twice its gradient, and which can overlap with the backward pass, so
1058it tolerates the slower network.
1059
1060**Will a 70-billion-parameter model serve from one 80 GB GPU?**
1061Not in 16-bit: the weights alone are 140 GB. In 8-bit the weights take 70
1062GB, leaving about 10 GB, around 30,000 tokens of 16-bit KV cache for a
1063model with 80 layers and 8 KV heads of 128 dimensions. In 4-bit, 45 GB is
1064left, about 137,000 tokens. Leave headroom for activations and the serving
1065software.
1066
1067## The papers behind this lesson
1068
1069- **Williams, Waterman and Patterson, *Roofline: An Insightful Visual
1070  Performance Model for Multicore Architectures* (2009)**:
1071  https://doi.org/10.1145/1498765.1498785. Introduced the roofline: judge a
1072  kernel by its arithmetic intensity against the machine's balance of
1073  compute and bandwidth.
1074- **Dao et al., *FlashAttention: Fast and Memory-Efficient Exact Attention
1075  with IO-Awareness* (2022)**: https://arxiv.org/abs/2205.14135. Applied
1076  tiling to attention so the score matrix never reaches HBM, showing that
1077  counting memory traffic rather than FLOPs is what makes attention fast.
1078  [Annotated companion](../../papers/flashattention.html)
1079- **Micikevicius et al., *Mixed Precision Training* (2017)**:
1080  https://arxiv.org/abs/1710.03740. Showed that networks train in 16-bit
1081  floats with fp32 master weights, fp32 accumulation and loss scaling.
1082- **Kalamkar et al., *A Study of BFLOAT16 for Deep Learning Training*
1083  (2019)**: https://arxiv.org/abs/1905.12322. Showed that bf16, with fp32's
1084  range, trains a wide range of models to fp32 quality without loss scaling.
1085- **Micikevicius et al., *FP8 Formats for Deep Learning* (2022)**:
1086  https://arxiv.org/abs/2209.05433. Proposed the E4M3 and E5M2 8-bit
1087  formats and showed that training and inference hold up in them.
1088- **Shoeybi et al., *Megatron-LM: Training Multi-Billion Parameter Language
1089  Models Using Model Parallelism* (2019)**: https://arxiv.org/abs/1909.08053.
1090  Split each transformer layer's matrix multiplies across the GPUs of one
1091  machine, the tensor parallelism of section 5.
1092  [Annotated companion](../../papers/megatron-lm.html)
1093- **Sergeev and Del Balso, *Horovod: fast and easy distributed deep
1094  learning in TensorFlow* (2018)**: https://arxiv.org/abs/1802.05799.
1095  Brought the bandwidth-optimal ring all-reduce to deep learning training.
1096
1097## Further reading
1098
1099- Horace He, *Making Deep Learning Go Brrrr From First Principles*: https://horace.io/brrr_intro.html
1100- *How to Scale Your Model* (a book on TPUs, GPUs and parallelism for transformers): https://jax-ml.github.io/scaling-book/
1101- Williams, Waterman and Patterson, *Roofline* (2009): https://doi.org/10.1145/1498765.1498785
1102- Dao et al., *FlashAttention* (2022): https://arxiv.org/abs/2205.14135
1103- Micikevicius et al., *Mixed Precision Training* (2017): https://arxiv.org/abs/1710.03740
1104- Micikevicius et al., *FP8 Formats for Deep Learning* (2022): https://arxiv.org/abs/2209.05433
1105- PyTorch automatic mixed precision: https://pytorch.org/docs/stable/amp.html
1106- NVIDIA, *CUDA C++ Programming Guide* (how one GPU family organises cores and memory): https://docs.nvidia.com/cuda/cuda-programming-guide/index.html
1107"""
1108
1109from __future__ import annotations
1110
1111import math
1112from dataclasses import dataclass
1113
1114import numpy as np
1115
1116from primer._show import banner, say, table, takeaway
1117from primer.ml.inference import HBM_BANDWIDTH, PEAK_FLOPS, kv_cache_bytes_per_token, weight_bytes
1118
1119# ---------------------------------------------------------------------------
1120# 1. The workload: a matrix multiply, and how it splits across cores
1121# ---------------------------------------------------------------------------
1122
1123
1124def matmul_flops(m: int, n: int, k: int) -> int:
1125    """FLOPs to multiply an (m, k) matrix by a (k, n) matrix.
1126
1127    Each of the m·n outputs is a dot product of length k: k multiplies and k
1128    adds, counted as 2 FLOPs per multiply-add (one fused instruction on a GPU).
1129    """
1130    return 2 * m * n * k
1131
1132
1133def counted_matmul(A: list[list[float]], B: list[list[float]]) -> tuple[list[list[float]], int, int]:
1134    """Multiply two small matrices with plain loops, counting every multiply and add.
1135
1136    Slow on purpose: the loop *is* the definition, and the counters prove the
1137    2·m·n·k rule instead of asserting it.
1138    """
1139    m, k, n = len(A), len(B), len(B[0])
1140    C = [[0 for _ in range(n)] for _ in range(m)]
1141    multiplies = adds = 0
1142    for i in range(m):
1143        for j in range(n):
1144            # One output cell: its own dot product, needing nothing from any other cell.
1145            total = 0
1146            for p in range(k):
1147                product = A[i][p] * B[p][j]
1148                multiplies += 1
1149                total = total + product
1150                adds += 1
1151            C[i][j] = total
1152    return C, multiplies, adds
1153
1154
1155def parallel_rounds(m: int, n: int, k: int, cores: int) -> int:
1156    """Rounds of multiply-adds to finish an (m,k)@(k,n) multiply on `cores` identical cores.
1157
1158    A deliberately simple model: each core takes whole output cells and works
1159    through a cell's k multiply-adds one per round. Cells are independent, so
1160    they share out perfectly until every cell has its own core; after that the
1161    extra cores have nothing to do and k rounds remain.
1162    """
1163    return math.ceil(m * n / cores) * k
1164
1165
1166# ---------------------------------------------------------------------------
1167# 2. The memory hierarchy
1168# ---------------------------------------------------------------------------
1169
1170
1171@dataclass(frozen=True)
1172class MemoryLevel:
1173    """One level of storage, as seen from the arithmetic units of one GPU."""
1174
1175    name: str
1176    capacity_bytes: float
1177    bandwidth_bytes_per_s: float
1178    picture: str  # the kitchen analogy used in the lesson
1179
1180
1181# Round, illustrative orders of magnitude for one datacenter GPU and its
1182# machine, not any product's specification. HBM uses the same bandwidth as
1183# primer.ml.inference so the two lessons describe the same imaginary chip.
1184MEMORY_HIERARCHY: list[MemoryLevel] = [
1185    MemoryLevel("registers", 20e6, 100e12, "the cook's hands"),
1186    MemoryLevel("on-chip SRAM", 50e6, 20e12, "the cutting board"),
1187    MemoryLevel("HBM", 80e9, HBM_BANDWIDTH, "the fridge in the kitchen"),
1188    MemoryLevel("host memory", 1e12, 50e9, "the storeroom down the hall"),
1189    MemoryLevel("local disk", 10e12, 10e9, "the warehouse across town"),
1190    MemoryLevel("network", 1e15, 50e9, "other restaurants' pantries, by courier"),
1191]
1192
1193
1194def level(name: str) -> MemoryLevel:
1195    """Look up a level of `MEMORY_HIERARCHY` by name."""
1196    return next(lv for lv in MEMORY_HIERARCHY if lv.name == name)
1197
1198
1199def transfer_seconds(nbytes: float, level_name: str) -> float:
1200    """Time to stream `nbytes` from one level at its full bandwidth (ignoring latency)."""
1201    return nbytes / level(level_name).bandwidth_bytes_per_s
1202
1203
1204# ---------------------------------------------------------------------------
1205# 3. Arithmetic intensity: a tiled matrix multiply that counts its traffic
1206# ---------------------------------------------------------------------------
1207
1208
1209@dataclass
1210class Traffic:
1211    """What a kernel moved between slow memory (HBM) and fast memory (SRAM), in elements."""
1212
1213    reads: int = 0
1214    writes: int = 0
1215    flops: int = 0
1216    fast_memory_peak: int = 0
1217
1218
1219def tiled_matmul(A: np.ndarray, B: np.ndarray, tile: int) -> tuple[np.ndarray, Traffic]:
1220    """C = A @ B computed tile by tile, counting every element that crosses from slow to fast memory.
1221
1222    For each (tile × tile) block of C: keep an accumulator in fast memory, and
1223    walk along the shared dimension loading one tile of A and one tile of B at
1224    a time. Each loaded element is then used `tile` times before it is thrown
1225    away, which is the reuse that cuts traffic. `tile=1` is no reuse at all:
1226    every multiply fetches both of its inputs.
1227    """
1228    m, k = A.shape
1229    k2, n = B.shape
1230    assert k == k2, "inner dimensions must match"
1231    assert m % tile == n % tile == k % tile == 0, "tiles must divide the matrix evenly"
1232    C = np.zeros((m, n))
1233    t = Traffic()
1234    for i0 in range(0, m, tile):
1235        for j0 in range(0, n, tile):
1236            acc = np.zeros((tile, tile))  # lives in fast memory for the whole walk along k
1237            for k0 in range(0, k, tile):
1238                a = A[i0:i0 + tile, k0:k0 + tile]  # load from slow memory
1239                b = B[k0:k0 + tile, j0:j0 + tile]
1240                t.reads += a.size + b.size
1241                t.fast_memory_peak = max(t.fast_memory_peak, a.size + b.size + acc.size)
1242                acc += a @ b  # tile³ multiply-adds on data already on chip
1243                t.flops += 2 * a.shape[0] * b.shape[1] * a.shape[1]
1244            C[i0:i0 + tile, j0:j0 + tile] = acc  # one write per output element, at the end
1245            t.writes += acc.size
1246    return C, t
1247
1248
1249def matmul_reads(n: int, tile: int) -> int:
1250    """Elements read from slow memory by `tiled_matmul` for two (n, n) matrices: 2n³ / tile."""
1251    return 2 * n**3 // tile
1252
1253
1254def matmul_intensity(n: int, tile: int, bytes_per_element: float) -> float:
1255    """FLOPs per byte of slow-memory traffic for an (n, n) tiled multiply, reads plus writes."""
1256    moved = (matmul_reads(n, tile) + n * n) * bytes_per_element
1257    return matmul_flops(n, n, n) / moved
1258
1259
1260# ---------------------------------------------------------------------------
1261# 4. Number formats, from scratch
1262# ---------------------------------------------------------------------------
1263
1264
1265@dataclass(frozen=True)
1266class FloatFormat:
1267    """A binary floating-point format: 1 sign bit, some exponent bits, some mantissa bits.
1268
1269    A normal value is (-1)^sign × 2^(exponent - bias) × (1 + mantissa / 2^M).
1270    `has_infinity` formats (IEEE style) reserve the all-ones exponent for
1271    infinity and NaN. The fp8 E4M3 format does not: it keeps that exponent
1272    for ordinary numbers and gives up only one pattern, to NaN, which buys it
1273    one more power of two of range.
1274    """
1275
1276    name: str
1277    exponent_bits: int
1278    mantissa_bits: int
1279    has_infinity: bool = True
1280
1281    @property
1282    def bits(self) -> int:
1283        return 1 + self.exponent_bits + self.mantissa_bits
1284
1285    @property
1286    def bias(self) -> int:
1287        # Stored exponents are unsigned; the bias centres them so small and large numbers get equal room.
1288        return 2 ** (self.exponent_bits - 1) - 1
1289
1290    @property
1291    def max_exponent_field(self) -> int:
1292        top = 2**self.exponent_bits - 1
1293        return top - 1 if self.has_infinity else top
1294
1295    @property
1296    def max_mantissa_at_top(self) -> int:
1297        # E4M3 spends the all-ones mantissa at the top exponent on NaN.
1298        return 2**self.mantissa_bits - 1 if self.has_infinity else 2**self.mantissa_bits - 2
1299
1300    @property
1301    def max_value(self) -> float:
1302        """The largest finite number the format can hold."""
1303        return (1 + self.max_mantissa_at_top / 2**self.mantissa_bits) * 2.0 ** (self.max_exponent_field - self.bias)
1304
1305    @property
1306    def min_normal(self) -> float:
1307        """The smallest number held at full precision."""
1308        return 2.0 ** (1 - self.bias)
1309
1310    @property
1311    def min_subnormal(self) -> float:
1312        """The smallest number above zero at all (with a single significant bit)."""
1313        return 2.0 ** (1 - self.bias - self.mantissa_bits)
1314
1315    @property
1316    def epsilon(self) -> float:
1317        """The gap between 1 and the next number up: the format's relative precision."""
1318        return 2.0**-self.mantissa_bits
1319
1320
1321FP32 = FloatFormat("fp32", 8, 23)
1322BF16 = FloatFormat("bf16", 8, 7)
1323FP16 = FloatFormat("fp16", 5, 10)
1324FP8_E5M2 = FloatFormat("fp8 E5M2", 5, 2)
1325FP8_E4M3 = FloatFormat("fp8 E4M3", 4, 3, has_infinity=False)
1326FLOAT_FORMATS = [FP32, BF16, FP16, FP8_E5M2, FP8_E4M3]
1327
1328
1329def _overflow_fields(sign: int, fmt: FloatFormat) -> tuple[int, int, int]:
1330    """Fields for a value too big to hold: infinity if the format has it, otherwise its NaN."""
1331    if fmt.has_infinity:
1332        return sign, 2**fmt.exponent_bits - 1, 0
1333    return sign, 2**fmt.exponent_bits - 1, 2**fmt.mantissa_bits - 1
1334
1335
1336def encode(x: float, fmt: FloatFormat) -> tuple[int, int, int]:
1337    """Round `x` to the nearest value `fmt` can hold and return its (sign, exponent, mantissa) fields.
1338
1339    Plain arithmetic, no bit tricks: find the power of two just below |x|,
1340    count how many steps of that binade's spacing |x| is, round to a whole
1341    number of steps (ties to even, like hardware), and split the step count
1342    into the hidden leading 1 and the stored mantissa.
1343    """
1344    M, bias = fmt.mantissa_bits, fmt.bias
1345    if math.isnan(x):
1346        return 0, 2**fmt.exponent_bits - 1, 2**M - 1
1347    sign = 1 if math.copysign(1.0, x) < 0 else 0
1348    a = abs(x)
1349    if a == 0.0:
1350        return sign, 0, 0
1351    if math.isinf(a):
1352        return _overflow_fields(sign, fmt)
1353    # frexp gives a = f · 2^e with f in [0.5, 1), so a lies in [2^(e-1), 2^e).
1354    exp = math.frexp(a)[1] - 1
1355    # Below the normal range the spacing stops shrinking: that is what subnormals are.
1356    exp = max(exp, 1 - bias)
1357    # a measured in steps of 2^(exp - M); Python's round() breaks ties to even, as hardware does.
1358    steps = round(a / 2.0 ** (exp - M))
1359    if steps == 2 ** (M + 1):
1360        # Rounding carried into the next power of two: 1.111… became 10.000….
1361        steps, exp = steps // 2, exp + 1
1362    if steps < 2**M:
1363        field, mantissa = 0, steps  # subnormal (or zero): no hidden leading 1
1364    else:
1365        field, mantissa = exp + bias, steps - 2**M  # the leading 1 is implied, not stored
1366    if field > fmt.max_exponent_field or (field == fmt.max_exponent_field and mantissa > fmt.max_mantissa_at_top):
1367        return _overflow_fields(sign, fmt)
1368    return sign, field, mantissa
1369
1370
1371def decode(sign: int, exponent: int, mantissa: int, fmt: FloatFormat) -> float:
1372    """The number a (sign, exponent, mantissa) triple stands for in `fmt`."""
1373    M, bias, top = fmt.mantissa_bits, fmt.bias, 2**fmt.exponent_bits - 1
1374    if fmt.has_infinity and exponent == top:
1375        magnitude = math.inf if mantissa == 0 else math.nan
1376    elif not fmt.has_infinity and exponent == top and mantissa == 2**M - 1:
1377        magnitude = math.nan
1378    elif exponent == 0:
1379        magnitude = mantissa * 2.0 ** (1 - bias - M)  # subnormal: 0.mantissa × 2^(1 - bias)
1380    else:
1381        magnitude = (2**M + mantissa) * 2.0 ** (exponent - bias - M)  # 1.mantissa × 2^(exponent - bias)
1382    return -magnitude if sign else magnitude
1383
1384
1385def round_to(x: float, fmt: FloatFormat) -> float:
1386    """`x` as it comes back after being stored in `fmt`."""
1387    return decode(*encode(x, fmt), fmt)
1388
1389
1390def bit_string(x: float, fmt: FloatFormat) -> str:
1391    """The stored bits of `x`, spaced as sign, exponent, mantissa."""
1392    s, e, m = encode(x, fmt)
1393    return f"{s} {e:0{fmt.exponent_bits}b} {m:0{fmt.mantissa_bits}b}"
1394
1395
1396def multiplier_cells(fmt: FloatFormat) -> int:
1397    """One-bit cells in a schoolbook (array) multiplier for the format's significands.
1398
1399    Multiplying two p-bit numbers the long-multiplication way needs p × p
1400    one-bit products, with p = mantissa bits + the hidden leading 1. The
1401    exponents only need adding, which is cheap by comparison.
1402    """
1403    p = fmt.mantissa_bits + 1
1404    return p * p
1405
1406
1407# ---------------------------------------------------------------------------
1408# 5. Many GPUs: ring all-reduce and the cost of talking
1409# ---------------------------------------------------------------------------
1410
1411# Illustrative, per GPU: a fast link between GPUs inside one machine, and the
1412# network card each GPU uses to reach GPUs in other machines (the "network"
1413# row of MEMORY_HIERARCHY).
1414IN_MACHINE_LINK = 500e9
1415BETWEEN_MACHINES_LINK = level("network").bandwidth_bytes_per_s
1416
1417
1418def ring_all_reduce(vectors: list[np.ndarray]) -> tuple[list[np.ndarray], list[int]]:
1419    """Sum one vector per GPU so every GPU ends with the total, passing chunks around a ring.
1420
1421    Each GPU splits its vector into N chunks. Phase 1 (reduce-scatter): for
1422    N-1 steps, every GPU passes one chunk to its right-hand neighbour, which
1423    adds it to its own copy; afterwards GPU g owns the finished sum of chunk
1424    (g + 1) mod N. Phase 2 (all-gather): for N-1 more steps, the finished
1425    chunks travel round the ring and overwrite the stale copies.
1426
1427    Returns every GPU's final vector and how many numbers each GPU sent.
1428    """
1429    N = len(vectors)
1430    chunks = [np.array_split(v.astype(float).copy(), N) for v in vectors]
1431    sent = [0] * N
1432    for step in range(N - 1):
1433        # All sends in a step happen at once, so read every outgoing chunk before anyone adds.
1434        outgoing = [(g, (g - step) % N, chunks[g][(g - step) % N].copy()) for g in range(N)]
1435        for g, c, data in outgoing:
1436            chunks[(g + 1) % N][c] += data
1437            sent[g] += data.size
1438    for step in range(N - 1):
1439        outgoing = [(g, (g + 1 - step) % N, chunks[g][(g + 1 - step) % N].copy()) for g in range(N)]
1440        for g, c, data in outgoing:
1441            chunks[(g + 1) % N][c] = data
1442            sent[g] += data.size
1443    return [np.concatenate(c) for c in chunks], sent
1444
1445
1446def all_reduce_seconds(nbytes: float, n_gpus: int, link_bytes_per_s: float) -> float:
1447    """Time for a ring all-reduce: each GPU sends 2(N-1)/N of the data over its link."""
1448    return 2 * (n_gpus - 1) / n_gpus * nbytes / link_bytes_per_s
1449
1450
1451def training_step_seconds(params: float, tokens: int, flops_per_second: float = PEAK_FLOPS) -> float:
1452    """Lower bound on one GPU's arithmetic for a training step: 6 FLOPs per parameter per token."""
1453    return 6 * params * tokens / flops_per_second
1454
1455
1456# ---------------------------------------------------------------------------
1457# 6. Serving: weights plus KV cache must fit
1458# ---------------------------------------------------------------------------
1459
1460# Llama 3 70B's attention shape: 80 layers, 8 key/value heads of 128 dimensions.
1461LLAMA3_70B_SHAPE = dict(layers=80, kv_heads=8, head_dim=128)
1462
1463
1464def max_cache_tokens(gpu_bytes: float, params: float, weight_bits: int, kv_bits: int,
1465                     layers: int, kv_heads: int, head_dim: int) -> int:
1466    """Tokens of KV cache that fit beside the weights (0 if the weights alone don't fit).
1467
1468    The formulas are primer.ml.inference's; this lesson varies the number format.
1469    """
1470    free = gpu_bytes - weight_bytes(params, weight_bits)
1471    if free <= 0:
1472        return 0
1473    return int(free // kv_cache_bytes_per_token(layers, kv_heads, head_dim, kv_bits))
1474
1475
1476# ---------------------------------------------------------------------------
1477# 7. Figures (rendered into the HTML docs by `make figures`)
1478# ---------------------------------------------------------------------------
1479
1480
1481def _relative_spacing(x: np.ndarray, fmt: FloatFormat) -> np.ndarray:
1482    """Gap to the next storable value, divided by x; NaN past the format's largest value."""
1483    exp = np.maximum(np.floor(np.log2(x)), 1 - fmt.bias)  # subnormals share the smallest normal spacing
1484    rel = 2.0 ** (exp - fmt.mantissa_bits) / x
1485    return np.where(x <= fmt.max_value, rel, np.nan)
1486
1487
1488def figures() -> dict:
1489    """Plot this lesson's data. matplotlib is imported here, and only here,
1490    so the lesson itself needs nothing beyond NumPy."""
1491    import matplotlib
1492
1493    matplotlib.use("Agg")
1494    import matplotlib.pyplot as plt
1495
1496    BLUE, RED, GREEN, ORANGE, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1497    figs = {}
1498
1499    # --- 1. Parallel rounds: speed-up until the independent work runs out ----
1500    cores = 2 ** np.arange(0, 15)
1501    fig, ax = plt.subplots(figsize=(6, 3.4))
1502    ax.plot(cores, [parallel_rounds(64, 64, 64, int(c)) for c in cores], "o-", color=BLUE)
1503    ax.set_xscale("log", base=2)
1504    ax.set_yscale("log")
1505    ax.axvline(64 * 64, color=MUTED, ls="--")
1506    ax.text(64 * 64 / 1.3, 1e4, "cores = cells\n(4,096)", color="#4b5563", ha="right")
1507    ax.set(xlabel="cores working at once", ylabel="rounds to finish",
1508           title="A 64 × 64 × 64 multiply: faster until every cell has a core")
1509    figs["parallel_rounds"] = fig
1510
1511    # --- 2. The memory hierarchy: capacity up, bandwidth down --------------
1512    names = [lv.name for lv in MEMORY_HIERARCHY]
1513    y = np.arange(len(names))[::-1]
1514    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4), sharey=True)
1515    a1.barh(y, [lv.capacity_bytes for lv in MEMORY_HIERARCHY], color=BLUE)
1516    a2.barh(y, [lv.bandwidth_bytes_per_s for lv in MEMORY_HIERARCHY], color=ORANGE)
1517    a1.set_yticks(y, names)
1518    a1.set(xscale="log", xlabel="capacity (bytes)", title="Further out holds more...")
1519    a2.set(xscale="log", xlabel="bandwidth (bytes per second)", title="...and moves it more slowly")
1520    for a in (a1, a2):
1521        a.grid(axis="y", visible=False)
1522    fig.text(0.5, -0.02, "Illustrative orders of magnitude for one datacenter GPU, not a product specification",
1523             ha="center", color="#4b5563", fontsize=8)
1524    fig.tight_layout()
1525    figs["hierarchy"] = fig
1526
1527    # --- 3. Tiling: traffic falls, intensity rises --------------------------
1528    n_sim = 32
1529    tiles = [1, 2, 4, 8, 16, 32]
1530    measured = [tiled_matmul(np.ones((n_sim, n_sim)), np.ones((n_sim, n_sim)), t)[1].reads for t in tiles]
1531    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1532    a1.loglog(tiles, [matmul_reads(n_sim, t) for t in tiles], color=MUTED, base=2, label="formula 2n³/T")
1533    a1.loglog(tiles, measured, "o", color=BLUE, base=2, label="counted by tiled_matmul")
1534    a1.set(xlabel="tile width T", ylabel="numbers read from HBM", title=f"Reads for two {n_sim} × {n_sim} matrices")
1535    a1.legend(frameon=False)
1536    big = 4096
1537    ts = 2 ** np.arange(0, 12)
1538    for b, color, label in ((2, BLUE, "16-bit"), (1, GREEN, "8-bit")):
1539        a2.plot(ts, [matmul_intensity(big, int(t), b) for t in ts], "o-", color=color, label=label)
1540    a2.set_xscale("log", base=2)
1541    a2.set_yscale("log")
1542    ridge = PEAK_FLOPS / HBM_BANDWIDTH
1543    a2.axhline(ridge, color=RED, ls="--")
1544    a2.text(1.2, ridge * 1.25, f"break-even ≈ {ridge:.0f} FLOPs/byte", color=RED)
1545    a2.set(xlabel="tile width T", ylabel="FLOPs per byte moved", title=f"Intensity of a {big} × {big} multiply")
1546    a2.legend(frameon=False, loc="lower right")
1547    fig.tight_layout()
1548    figs["tiling"] = fig
1549
1550    # --- 4. Bit layouts, drawn to scale -------------------------------------
1551    rows = [(f.name, [("sign", 1), ("exponent", f.exponent_bits), ("mantissa", f.mantissa_bits)]) for f in FLOAT_FORMATS]
1552    rows += [("int8", [("integer", 8)]), ("int4", [("integer", 4)])]
1553    colors = {"sign": MUTED, "exponent": ORANGE, "mantissa": BLUE, "integer": GREEN}
1554    fig, ax = plt.subplots(figsize=(9, 3.6))
1555    for r, (name, parts) in enumerate(rows):
1556        yy, x0 = len(rows) - 1 - r, 0
1557        for part, width in parts:
1558            ax.barh(yy, width, left=x0, height=0.7, color=colors[part], edgecolor="white", linewidth=0.5)
1559            if width >= 2:
1560                ax.text(x0 + width / 2, yy, str(width), ha="center", va="center", color="white", fontsize=8)
1561            x0 += width
1562        ax.text(x0 + 0.4, yy, f"{x0} bits", va="center", fontsize=8, color="#4b5563")
1563    ax.set_yticks(range(len(rows)), [name for name, _ in rows][::-1])
1564    ax.set_xlim(0, 36)
1565    ax.set_xlabel("bits")
1566    ax.grid(False)
1567    handles = [plt.Rectangle((0, 0), 1, 1, color=c) for c in colors.values()]
1568    ax.legend(handles, list(colors), frameon=False, loc="lower right")
1569    ax.set_title("Where each format spends its bits")
1570    figs["format_layouts"] = fig
1571
1572    # --- 5. Relative precision across magnitudes ----------------------------
1573    x = np.logspace(-10, 6, 1200)
1574    fig, ax = plt.subplots(figsize=(8.5, 4))
1575    top = 10  # a line that would climb past the chart ends at its top instead of running into the title
1576
1577    def within(y):
1578        return np.where(y <= top, y, np.nan)
1579
1580    for fmt, color in zip(FLOAT_FORMATS, (BLUE, GREEN, ORANGE, RED, "#7c3aed")):
1581        ax.loglog(x, within(_relative_spacing(x, fmt)), color=color, label=fmt.name)
1582    for bits, ls in ((8, "--"), (4, ":")):
1583        step = 1 / (2 ** (bits - 1) - 1)  # scale chosen so 1.0 lands on the top code
1584        ax.loglog(x, within(np.where(x <= 1, step / x, np.nan)), color="#4b5563", ls=ls, label=f"int{bits}, scale 1/{2 ** (bits - 1) - 1}")
1585    ax.set_ylim(1e-8, top)
1586    ax.set(xlabel="size of the number stored", ylabel="gap to the next value ÷ the number",
1587           title="Relative precision: flat for floats, rising for integers")
1588    ax.legend(frameon=False, fontsize=8, loc="center left", bbox_to_anchor=(1.01, 0.5))
1589    fig.tight_layout()
1590    figs["precision"] = fig
1591
1592    # --- 6. Communication vs. computation -----------------------------------
1593    gpus = 2 ** np.arange(1, 11)
1594    grad_bytes, params, tokens = 14e9, 7e9, 16_384
1595    fig, ax = plt.subplots(figsize=(6.5, 3.6))
1596    ax.semilogx(gpus, [all_reduce_seconds(grad_bytes, int(g), BETWEEN_MACHINES_LINK) for g in gpus], "o-", color=RED,
1597                base=2, label="all-reduce over the network (50 GB/s)")
1598    ax.semilogx(gpus, [all_reduce_seconds(grad_bytes, int(g), IN_MACHINE_LINK) for g in gpus], "o-", color=BLUE,
1599                base=2, label="all-reduce over in-machine links (500 GB/s)")
1600    ax.axhline(training_step_seconds(params, tokens), color=MUTED, lw=3, label="arithmetic per step (16k tokens)")
1601    ax.set_ylim(0, 0.8)
1602    ax.set(xlabel="GPUs in the all-reduce", ylabel="seconds per training step",
1603           title="7B parameters, 16-bit gradients: talking vs. computing")
1604    ax.legend(frameon=False, fontsize=8, loc="center right")
1605    figs["communication"] = fig
1606
1607    # --- 7. Will it fit? ----------------------------------------------------
1608    gpu, params = 80e9, 70e9
1609    choices = [("16-bit", 16), ("fp8", 8), ("4-bit", 4)]
1610    fig, ax = plt.subplots(figsize=(6, 3.6))
1611    for i, (label, bits) in enumerate(choices):
1612        w = weight_bytes(params, bits) / 1e9
1613        free = max(gpu / 1e9 - w, 0)
1614        ax.bar(i, w, color=BLUE, label="weights" if i == 0 else None)
1615        ax.bar(i, free, bottom=w, color="#93c5fd", label="room for KV cache" if i == 0 else None)
1616        tokens_fit = max_cache_tokens(gpu, params, bits, 16, **LLAMA3_70B_SHAPE)
1617        note = "does not fit" if tokens_fit == 0 else f"{tokens_fit:,} tokens"
1618        ax.text(i, max(w, gpu / 1e9) + 3, note, ha="center", fontsize=9)
1619    ax.axhline(gpu / 1e9, color=RED, ls="--", label="80 GB of GPU memory")
1620    ax.set_xticks(range(len(choices)), [f"{label} weights" for label, _ in choices])
1621    ax.set_ylim(0, 185)
1622    ax.set(ylabel="GB", title="A 70B model on one GPU (16-bit KV cache)")
1623    ax.legend(frameon=False, loc="upper right")
1624    figs["serving_fit"] = fig
1625
1626    return figs
1627
1628
1629# ---------------------------------------------------------------------------
1630# 8. Narrated walkthrough
1631# ---------------------------------------------------------------------------
1632
1633
1634def demo() -> None:
1635    banner("1. The workload: a matrix multiply, counted by hand")
1636    C, muls, adds = counted_matmul([[1, 2, 3], [4, 5, 6]], [[7, 8], [9, 10], [11, 12]])
1637    say(f"[[1,2,3],[4,5,6]] times [[7,8],[9,10],[11,12]] = {C}.")
1638    say(
1639        f"""
1640        That took {muls} multiplies and {adds} adds: 2·m·n·k = {matmul_flops(2, 2, 3)} FLOPs.
1641        No output cell needed another cell's answer, so they could all be
1642        computed at once. A 4096-cube multiply is {matmul_flops(4096, 4096, 4096):.3e} FLOPs.
1643        """
1644    )
1645    table(["cores", "rounds for a 64 × 64 × 64 multiply"], [(c, parallel_rounds(64, 64, 64, c)) for c in (1, 64, 4096, 8192)])
1646    takeaway("A GPU wins by doing thousands of independent multiply-adds at once; it needs that much independent work.")
1647
1648    banner("2. The memory hierarchy")
1649    table(
1650        ["level", "holds", "moves per second", "ms to stream 1 GB", "picture"],
1651        [(lv.name, f"{lv.capacity_bytes:.0e} B", f"{lv.bandwidth_bytes_per_s:.2e} B", transfer_seconds(1e9, lv.name) * 1e3, lv.picture)
1652         for lv in MEMORY_HIERARCHY],
1653        floatfmt=".2f",
1654    )
1655    say(
1656        """
1657        Illustrative round numbers. Reading 16 GB of weights from HBM takes
1658        4.78 ms; from host memory it would take 320 ms.
1659        """
1660    )
1661    takeaway("Each step away from the arithmetic is bigger and slower; full speed means living in HBM or closer.")
1662
1663    banner("3. Tiling: same FLOPs, far fewer bytes")
1664    rows = []
1665    for t in (1, 2, 4):
1666        _, traffic = tiled_matmul(np.ones((4, 4)), np.ones((4, 4)), t)
1667        rows.append((f"{t} × {t}", traffic.reads, traffic.writes, traffic.flops, traffic.flops / (traffic.reads + traffic.writes)))
1668    table(["tile", "reads", "writes", "FLOPs", "FLOPs per number moved"], rows, floatfmt=".2f")
1669    say(
1670        f"""
1671        For 4096 × 4096 matrices with 128-wide tiles: {matmul_intensity(4096, 128, 2):.1f} FLOPs per byte
1672        in 16-bit, {matmul_intensity(4096, 128, 1):.1f} in 8-bit. The chip breaks even at
1673        {PEAK_FLOPS / HBM_BANDWIDTH:.0f}. FlashAttention, kernel fusion and batching all
1674        raise this ratio.
1675        """
1676    )
1677    takeaway("Speed is usually set by bytes moved, not FLOPs; reuse every byte you fetch as often as you can.")
1678
1679    banner("4. Number formats, built from scratch")
1680    say(f"-6.5 in bf16 is stored as {bit_string(-6.5, BF16)} (sign, exponent, mantissa).")
1681    table(
1682        ["format", "bits", "largest", "smallest normal", "gap above 1", "0.1 is stored as", "multiplier cells"],
1683        [(f.name, f.bits, f"{f.max_value:.4g}", f"{f.min_normal:.3g}", f"{f.epsilon:.3g}", f"{round_to(0.1, f):.8g}", multiplier_cells(f))
1684         for f in FLOAT_FORMATS],
1685    )
1686    say(
1687        f"""
1688        Range: 1e-8 in fp16 becomes {round_to(1e-8, FP16)}; in bf16 it is {round_to(1e-8, BF16):.4g}.
1689        Precision: 1 + 0.001 in bf16 is {round_to(1.001, BF16)}; in fp32 it is {round_to(1.001, FP32):.7g}.
1690        Overflow: 70,000 in fp16 is {round_to(70_000.0, FP16)}; 500 in fp8 E4M3 is {round_to(500.0, FP8_E4M3)}.
1691        """
1692    )
1693    takeaway("Exponent bits buy range, mantissa bits buy precision; fewer bits buy speed three ways.")
1694
1695    banner("5. Many GPUs: ring all-reduce")
1696    results, sent = ring_all_reduce([np.arange(8.0) * (g + 1) for g in range(4)])
1697    say(f"Four GPUs, 8 numbers each. After the ring, every GPU holds {results[0].tolist()}; each sent {sent[0]} numbers.")
1698    table(
1699        ["link", "all-reduce of 14 GB on 8 GPUs (ms)", "arithmetic per step (ms)"],
1700        [("in one machine", all_reduce_seconds(14e9, 8, IN_MACHINE_LINK) * 1e3, training_step_seconds(7e9, 16_384) * 1e3),
1701         ("between machines", all_reduce_seconds(14e9, 8, BETWEEN_MACHINES_LINK) * 1e3, training_step_seconds(7e9, 16_384) * 1e3)],
1702        floatfmt=".0f",
1703    )
1704    takeaway("Put chatty parallelism on fast links; send the once-per-step traffic over the network.")
1705
1706    banner("6. Will a 70B model fit on one 80 GB GPU?")
1707    table(
1708        ["weights", "weight GB", "tokens of 16-bit cache", "tokens of 8-bit cache"],
1709        [(f"{b}-bit", weight_bytes(70e9, b) / 1e9, max_cache_tokens(80e9, 70e9, b, 16, **LLAMA3_70B_SHAPE),
1710          max_cache_tokens(80e9, 70e9, b, 8, **LLAMA3_70B_SHAPE)) for b in (16, 8, 4)],
1711        floatfmt=".0f",
1712    )
1713    takeaway("The number format decides whether a model fits, how many users share a GPU, and how fast each token comes.")
1714
1715
1716if __name__ == "__main__":
1717    demo()
Level 3: the code, function by function.
def matmul_flops(m: int, n: int, k: int) -> int: on GitHub
1125def matmul_flops(m: int, n: int, k: int) -> int:
1126    """FLOPs to multiply an (m, k) matrix by a (k, n) matrix.
1127
1128    Each of the m·n outputs is a dot product of length k: k multiplies and k
1129    adds, counted as 2 FLOPs per multiply-add (one fused instruction on a GPU).
1130    """
1131    return 2 * m * n * k

FLOPs to multiply an (m, k) matrix by a (k, n) matrix.

Each of the m·n outputs is a dot product of length k: k multiplies and k adds, counted as 2 FLOPs per multiply-add (one fused instruction on a GPU).

def counted_matmul( A: list[list[float]], B: list[list[float]]) -> tuple[list[list[float]], int, int]: on GitHub
1134def counted_matmul(A: list[list[float]], B: list[list[float]]) -> tuple[list[list[float]], int, int]:
1135    """Multiply two small matrices with plain loops, counting every multiply and add.
1136
1137    Slow on purpose: the loop *is* the definition, and the counters prove the
1138    2·m·n·k rule instead of asserting it.
1139    """
1140    m, k, n = len(A), len(B), len(B[0])
1141    C = [[0 for _ in range(n)] for _ in range(m)]
1142    multiplies = adds = 0
1143    for i in range(m):
1144        for j in range(n):
1145            # One output cell: its own dot product, needing nothing from any other cell.
1146            total = 0
1147            for p in range(k):
1148                product = A[i][p] * B[p][j]
1149                multiplies += 1
1150                total = total + product
1151                adds += 1
1152            C[i][j] = total
1153    return C, multiplies, adds

Multiply two small matrices with plain loops, counting every multiply and add.

Slow on purpose: the loop is the definition, and the counters prove the 2·m·n·k rule instead of asserting it.

def parallel_rounds(m: int, n: int, k: int, cores: int) -> int: on GitHub
1156def parallel_rounds(m: int, n: int, k: int, cores: int) -> int:
1157    """Rounds of multiply-adds to finish an (m,k)@(k,n) multiply on `cores` identical cores.
1158
1159    A deliberately simple model: each core takes whole output cells and works
1160    through a cell's k multiply-adds one per round. Cells are independent, so
1161    they share out perfectly until every cell has its own core; after that the
1162    extra cores have nothing to do and k rounds remain.
1163    """
1164    return math.ceil(m * n / cores) * k

Rounds of multiply-adds to finish an (m,k)@(k,n) multiply on cores identical cores.

A deliberately simple model: each core takes whole output cells and works through a cell's k multiply-adds one per round. Cells are independent, so they share out perfectly until every cell has its own core; after that the extra cores have nothing to do and k rounds remain.

@dataclass(frozen=True)
class MemoryLevel: on GitHub
1172@dataclass(frozen=True)
1173class MemoryLevel:
1174    """One level of storage, as seen from the arithmetic units of one GPU."""
1175
1176    name: str
1177    capacity_bytes: float
1178    bandwidth_bytes_per_s: float
1179    picture: str  # the kitchen analogy used in the lesson

One level of storage, as seen from the arithmetic units of one GPU.

MemoryLevel( name: str, capacity_bytes: float, bandwidth_bytes_per_s: float, picture: str)
name: str
picture: str
MEMORY_HIERARCHY: list[MemoryLevel] = [MemoryLevel(name='registers', capacity_bytes=20000000.0, bandwidth_bytes_per_s=100000000000000.0, picture="the cook's hands"), MemoryLevel(name='on-chip SRAM', capacity_bytes=50000000.0, bandwidth_bytes_per_s=20000000000000.0, picture='the cutting board'), MemoryLevel(name='HBM', capacity_bytes=80000000000.0, bandwidth_bytes_per_s=3350000000000.0, picture='the fridge in the kitchen'), MemoryLevel(name='host memory', capacity_bytes=1000000000000.0, bandwidth_bytes_per_s=50000000000.0, picture='the storeroom down the hall'), MemoryLevel(name='local disk', capacity_bytes=10000000000000.0, bandwidth_bytes_per_s=10000000000.0, picture='the warehouse across town'), MemoryLevel(name='network', capacity_bytes=1000000000000000.0, bandwidth_bytes_per_s=50000000000.0, picture="other restaurants' pantries, by courier")]
def level(name: str) -> MemoryLevel: on GitHub
1195def level(name: str) -> MemoryLevel:
1196    """Look up a level of `MEMORY_HIERARCHY` by name."""
1197    return next(lv for lv in MEMORY_HIERARCHY if lv.name == name)

Look up a level of MEMORY_HIERARCHY by name.

def transfer_seconds(nbytes: float, level_name: str) -> float: on GitHub
1200def transfer_seconds(nbytes: float, level_name: str) -> float:
1201    """Time to stream `nbytes` from one level at its full bandwidth (ignoring latency)."""
1202    return nbytes / level(level_name).bandwidth_bytes_per_s

Time to stream nbytes from one level at its full bandwidth (ignoring latency).

@dataclass
class Traffic: on GitHub
1210@dataclass
1211class Traffic:
1212    """What a kernel moved between slow memory (HBM) and fast memory (SRAM), in elements."""
1213
1214    reads: int = 0
1215    writes: int = 0
1216    flops: int = 0
1217    fast_memory_peak: int = 0

What a kernel moved between slow memory (HBM) and fast memory (SRAM), in elements.

Traffic( reads: int = 0, writes: int = 0, flops: int = 0, fast_memory_peak: int = 0)
reads: int = 0
writes: int = 0
flops: int = 0
fast_memory_peak: int = 0
def tiled_matmul( A: numpy.ndarray, B: numpy.ndarray, tile: int) -> tuple[numpy.ndarray, Traffic]: on GitHub
1220def tiled_matmul(A: np.ndarray, B: np.ndarray, tile: int) -> tuple[np.ndarray, Traffic]:
1221    """C = A @ B computed tile by tile, counting every element that crosses from slow to fast memory.
1222
1223    For each (tile × tile) block of C: keep an accumulator in fast memory, and
1224    walk along the shared dimension loading one tile of A and one tile of B at
1225    a time. Each loaded element is then used `tile` times before it is thrown
1226    away, which is the reuse that cuts traffic. `tile=1` is no reuse at all:
1227    every multiply fetches both of its inputs.
1228    """
1229    m, k = A.shape
1230    k2, n = B.shape
1231    assert k == k2, "inner dimensions must match"
1232    assert m % tile == n % tile == k % tile == 0, "tiles must divide the matrix evenly"
1233    C = np.zeros((m, n))
1234    t = Traffic()
1235    for i0 in range(0, m, tile):
1236        for j0 in range(0, n, tile):
1237            acc = np.zeros((tile, tile))  # lives in fast memory for the whole walk along k
1238            for k0 in range(0, k, tile):
1239                a = A[i0:i0 + tile, k0:k0 + tile]  # load from slow memory
1240                b = B[k0:k0 + tile, j0:j0 + tile]
1241                t.reads += a.size + b.size
1242                t.fast_memory_peak = max(t.fast_memory_peak, a.size + b.size + acc.size)
1243                acc += a @ b  # tile³ multiply-adds on data already on chip
1244                t.flops += 2 * a.shape[0] * b.shape[1] * a.shape[1]
1245            C[i0:i0 + tile, j0:j0 + tile] = acc  # one write per output element, at the end
1246            t.writes += acc.size
1247    return C, t

C = A @ B computed tile by tile, counting every element that crosses from slow to fast memory.

For each (tile × tile) block of C: keep an accumulator in fast memory, and walk along the shared dimension loading one tile of A and one tile of B at a time. Each loaded element is then used tile times before it is thrown away, which is the reuse that cuts traffic. tile=1 is no reuse at all: every multiply fetches both of its inputs.

def matmul_reads(n: int, tile: int) -> int: on GitHub
1250def matmul_reads(n: int, tile: int) -> int:
1251    """Elements read from slow memory by `tiled_matmul` for two (n, n) matrices: 2n³ / tile."""
1252    return 2 * n**3 // tile

Elements read from slow memory by tiled_matmul for two (n, n) matrices: 2n³ / tile.

def matmul_intensity(n: int, tile: int, bytes_per_element: float) -> float: on GitHub
1255def matmul_intensity(n: int, tile: int, bytes_per_element: float) -> float:
1256    """FLOPs per byte of slow-memory traffic for an (n, n) tiled multiply, reads plus writes."""
1257    moved = (matmul_reads(n, tile) + n * n) * bytes_per_element
1258    return matmul_flops(n, n, n) / moved

FLOPs per byte of slow-memory traffic for an (n, n) tiled multiply, reads plus writes.

@dataclass(frozen=True)
class FloatFormat: on GitHub
1266@dataclass(frozen=True)
1267class FloatFormat:
1268    """A binary floating-point format: 1 sign bit, some exponent bits, some mantissa bits.
1269
1270    A normal value is (-1)^sign × 2^(exponent - bias) × (1 + mantissa / 2^M).
1271    `has_infinity` formats (IEEE style) reserve the all-ones exponent for
1272    infinity and NaN. The fp8 E4M3 format does not: it keeps that exponent
1273    for ordinary numbers and gives up only one pattern, to NaN, which buys it
1274    one more power of two of range.
1275    """
1276
1277    name: str
1278    exponent_bits: int
1279    mantissa_bits: int
1280    has_infinity: bool = True
1281
1282    @property
1283    def bits(self) -> int:
1284        return 1 + self.exponent_bits + self.mantissa_bits
1285
1286    @property
1287    def bias(self) -> int:
1288        # Stored exponents are unsigned; the bias centres them so small and large numbers get equal room.
1289        return 2 ** (self.exponent_bits - 1) - 1
1290
1291    @property
1292    def max_exponent_field(self) -> int:
1293        top = 2**self.exponent_bits - 1
1294        return top - 1 if self.has_infinity else top
1295
1296    @property
1297    def max_mantissa_at_top(self) -> int:
1298        # E4M3 spends the all-ones mantissa at the top exponent on NaN.
1299        return 2**self.mantissa_bits - 1 if self.has_infinity else 2**self.mantissa_bits - 2
1300
1301    @property
1302    def max_value(self) -> float:
1303        """The largest finite number the format can hold."""
1304        return (1 + self.max_mantissa_at_top / 2**self.mantissa_bits) * 2.0 ** (self.max_exponent_field - self.bias)
1305
1306    @property
1307    def min_normal(self) -> float:
1308        """The smallest number held at full precision."""
1309        return 2.0 ** (1 - self.bias)
1310
1311    @property
1312    def min_subnormal(self) -> float:
1313        """The smallest number above zero at all (with a single significant bit)."""
1314        return 2.0 ** (1 - self.bias - self.mantissa_bits)
1315
1316    @property
1317    def epsilon(self) -> float:
1318        """The gap between 1 and the next number up: the format's relative precision."""
1319        return 2.0**-self.mantissa_bits

A binary floating-point format: 1 sign bit, some exponent bits, some mantissa bits.

A normal value is (-1)^sign × 2^(exponent - bias) × (1 + mantissa / 2^M). has_infinity formats (IEEE style) reserve the all-ones exponent for infinity and NaN. The fp8 E4M3 format does not: it keeps that exponent for ordinary numbers and gives up only one pattern, to NaN, which buys it one more power of two of range.

FloatFormat( name: str, exponent_bits: int, mantissa_bits: int, has_infinity: bool = True)
name: str
has_infinity: bool = True
bits: int on GitHub
1282    @property
1283    def bits(self) -> int:
1284        return 1 + self.exponent_bits + self.mantissa_bits
bias: int on GitHub
1286    @property
1287    def bias(self) -> int:
1288        # Stored exponents are unsigned; the bias centres them so small and large numbers get equal room.
1289        return 2 ** (self.exponent_bits - 1) - 1
max_exponent_field: int on GitHub
1291    @property
1292    def max_exponent_field(self) -> int:
1293        top = 2**self.exponent_bits - 1
1294        return top - 1 if self.has_infinity else top
max_mantissa_at_top: int on GitHub
1296    @property
1297    def max_mantissa_at_top(self) -> int:
1298        # E4M3 spends the all-ones mantissa at the top exponent on NaN.
1299        return 2**self.mantissa_bits - 1 if self.has_infinity else 2**self.mantissa_bits - 2
max_value: float on GitHub
1301    @property
1302    def max_value(self) -> float:
1303        """The largest finite number the format can hold."""
1304        return (1 + self.max_mantissa_at_top / 2**self.mantissa_bits) * 2.0 ** (self.max_exponent_field - self.bias)

The largest finite number the format can hold.

min_normal: float on GitHub
1306    @property
1307    def min_normal(self) -> float:
1308        """The smallest number held at full precision."""
1309        return 2.0 ** (1 - self.bias)

The smallest number held at full precision.

min_subnormal: float on GitHub
1311    @property
1312    def min_subnormal(self) -> float:
1313        """The smallest number above zero at all (with a single significant bit)."""
1314        return 2.0 ** (1 - self.bias - self.mantissa_bits)

The smallest number above zero at all (with a single significant bit).

epsilon: float on GitHub
1316    @property
1317    def epsilon(self) -> float:
1318        """The gap between 1 and the next number up: the format's relative precision."""
1319        return 2.0**-self.mantissa_bits

The gap between 1 and the next number up: the format's relative precision.

FP32 = FloatFormat(name='fp32', exponent_bits=8, mantissa_bits=23, has_infinity=True)
BF16 = FloatFormat(name='bf16', exponent_bits=8, mantissa_bits=7, has_infinity=True)
FP16 = FloatFormat(name='fp16', exponent_bits=5, mantissa_bits=10, has_infinity=True)
FP8_E5M2 = FloatFormat(name='fp8 E5M2', exponent_bits=5, mantissa_bits=2, has_infinity=True)
FP8_E4M3 = FloatFormat(name='fp8 E4M3', exponent_bits=4, mantissa_bits=3, has_infinity=False)
FLOAT_FORMATS = [FloatFormat(name='fp32', exponent_bits=8, mantissa_bits=23, has_infinity=True), FloatFormat(name='bf16', exponent_bits=8, mantissa_bits=7, has_infinity=True), FloatFormat(name='fp16', exponent_bits=5, mantissa_bits=10, has_infinity=True), FloatFormat(name='fp8 E5M2', exponent_bits=5, mantissa_bits=2, has_infinity=True), FloatFormat(name='fp8 E4M3', exponent_bits=4, mantissa_bits=3, has_infinity=False)]
def encode(x: float, fmt: FloatFormat) -> tuple[int, int, int]: on GitHub
1337def encode(x: float, fmt: FloatFormat) -> tuple[int, int, int]:
1338    """Round `x` to the nearest value `fmt` can hold and return its (sign, exponent, mantissa) fields.
1339
1340    Plain arithmetic, no bit tricks: find the power of two just below |x|,
1341    count how many steps of that binade's spacing |x| is, round to a whole
1342    number of steps (ties to even, like hardware), and split the step count
1343    into the hidden leading 1 and the stored mantissa.
1344    """
1345    M, bias = fmt.mantissa_bits, fmt.bias
1346    if math.isnan(x):
1347        return 0, 2**fmt.exponent_bits - 1, 2**M - 1
1348    sign = 1 if math.copysign(1.0, x) < 0 else 0
1349    a = abs(x)
1350    if a == 0.0:
1351        return sign, 0, 0
1352    if math.isinf(a):
1353        return _overflow_fields(sign, fmt)
1354    # frexp gives a = f · 2^e with f in [0.5, 1), so a lies in [2^(e-1), 2^e).
1355    exp = math.frexp(a)[1] - 1
1356    # Below the normal range the spacing stops shrinking: that is what subnormals are.
1357    exp = max(exp, 1 - bias)
1358    # a measured in steps of 2^(exp - M); Python's round() breaks ties to even, as hardware does.
1359    steps = round(a / 2.0 ** (exp - M))
1360    if steps == 2 ** (M + 1):
1361        # Rounding carried into the next power of two: 1.111… became 10.000….
1362        steps, exp = steps // 2, exp + 1
1363    if steps < 2**M:
1364        field, mantissa = 0, steps  # subnormal (or zero): no hidden leading 1
1365    else:
1366        field, mantissa = exp + bias, steps - 2**M  # the leading 1 is implied, not stored
1367    if field > fmt.max_exponent_field or (field == fmt.max_exponent_field and mantissa > fmt.max_mantissa_at_top):
1368        return _overflow_fields(sign, fmt)
1369    return sign, field, mantissa

Round x to the nearest value fmt can hold and return its (sign, exponent, mantissa) fields.

Plain arithmetic, no bit tricks: find the power of two just below |x|, count how many steps of that binade's spacing |x| is, round to a whole number of steps (ties to even, like hardware), and split the step count into the hidden leading 1 and the stored mantissa.

def decode( sign: int, exponent: int, mantissa: int, fmt: FloatFormat) -> float: on GitHub
1372def decode(sign: int, exponent: int, mantissa: int, fmt: FloatFormat) -> float:
1373    """The number a (sign, exponent, mantissa) triple stands for in `fmt`."""
1374    M, bias, top = fmt.mantissa_bits, fmt.bias, 2**fmt.exponent_bits - 1
1375    if fmt.has_infinity and exponent == top:
1376        magnitude = math.inf if mantissa == 0 else math.nan
1377    elif not fmt.has_infinity and exponent == top and mantissa == 2**M - 1:
1378        magnitude = math.nan
1379    elif exponent == 0:
1380        magnitude = mantissa * 2.0 ** (1 - bias - M)  # subnormal: 0.mantissa × 2^(1 - bias)
1381    else:
1382        magnitude = (2**M + mantissa) * 2.0 ** (exponent - bias - M)  # 1.mantissa × 2^(exponent - bias)
1383    return -magnitude if sign else magnitude

The number a (sign, exponent, mantissa) triple stands for in fmt.

def round_to(x: float, fmt: FloatFormat) -> float: on GitHub
1386def round_to(x: float, fmt: FloatFormat) -> float:
1387    """`x` as it comes back after being stored in `fmt`."""
1388    return decode(*encode(x, fmt), fmt)

x as it comes back after being stored in fmt.

def bit_string(x: float, fmt: FloatFormat) -> str: on GitHub
1391def bit_string(x: float, fmt: FloatFormat) -> str:
1392    """The stored bits of `x`, spaced as sign, exponent, mantissa."""
1393    s, e, m = encode(x, fmt)
1394    return f"{s} {e:0{fmt.exponent_bits}b} {m:0{fmt.mantissa_bits}b}"

The stored bits of x, spaced as sign, exponent, mantissa.

def multiplier_cells(fmt: FloatFormat) -> int: on GitHub
1397def multiplier_cells(fmt: FloatFormat) -> int:
1398    """One-bit cells in a schoolbook (array) multiplier for the format's significands.
1399
1400    Multiplying two p-bit numbers the long-multiplication way needs p × p
1401    one-bit products, with p = mantissa bits + the hidden leading 1. The
1402    exponents only need adding, which is cheap by comparison.
1403    """
1404    p = fmt.mantissa_bits + 1
1405    return p * p

One-bit cells in a schoolbook (array) multiplier for the format's significands.

Multiplying two p-bit numbers the long-multiplication way needs p × p one-bit products, with p = mantissa bits + the hidden leading 1. The exponents only need adding, which is cheap by comparison.

def ring_all_reduce(vectors: list[numpy.ndarray]) -> tuple[list[numpy.ndarray], list[int]]: on GitHub
1419def ring_all_reduce(vectors: list[np.ndarray]) -> tuple[list[np.ndarray], list[int]]:
1420    """Sum one vector per GPU so every GPU ends with the total, passing chunks around a ring.
1421
1422    Each GPU splits its vector into N chunks. Phase 1 (reduce-scatter): for
1423    N-1 steps, every GPU passes one chunk to its right-hand neighbour, which
1424    adds it to its own copy; afterwards GPU g owns the finished sum of chunk
1425    (g + 1) mod N. Phase 2 (all-gather): for N-1 more steps, the finished
1426    chunks travel round the ring and overwrite the stale copies.
1427
1428    Returns every GPU's final vector and how many numbers each GPU sent.
1429    """
1430    N = len(vectors)
1431    chunks = [np.array_split(v.astype(float).copy(), N) for v in vectors]
1432    sent = [0] * N
1433    for step in range(N - 1):
1434        # All sends in a step happen at once, so read every outgoing chunk before anyone adds.
1435        outgoing = [(g, (g - step) % N, chunks[g][(g - step) % N].copy()) for g in range(N)]
1436        for g, c, data in outgoing:
1437            chunks[(g + 1) % N][c] += data
1438            sent[g] += data.size
1439    for step in range(N - 1):
1440        outgoing = [(g, (g + 1 - step) % N, chunks[g][(g + 1 - step) % N].copy()) for g in range(N)]
1441        for g, c, data in outgoing:
1442            chunks[(g + 1) % N][c] = data
1443            sent[g] += data.size
1444    return [np.concatenate(c) for c in chunks], sent

Sum one vector per GPU so every GPU ends with the total, passing chunks around a ring.

Each GPU splits its vector into N chunks. Phase 1 (reduce-scatter): for N-1 steps, every GPU passes one chunk to its right-hand neighbour, which adds it to its own copy; afterwards GPU g owns the finished sum of chunk (g + 1) mod N. Phase 2 (all-gather): for N-1 more steps, the finished chunks travel round the ring and overwrite the stale copies.

Returns every GPU's final vector and how many numbers each GPU sent.

def all_reduce_seconds(nbytes: float, n_gpus: int, link_bytes_per_s: float) -> float: on GitHub
1447def all_reduce_seconds(nbytes: float, n_gpus: int, link_bytes_per_s: float) -> float:
1448    """Time for a ring all-reduce: each GPU sends 2(N-1)/N of the data over its link."""
1449    return 2 * (n_gpus - 1) / n_gpus * nbytes / link_bytes_per_s

Time for a ring all-reduce: each GPU sends 2(N-1)/N of the data over its link.

def training_step_seconds( params: float, tokens: int, flops_per_second: float = 1000000000000000.0) -> float: on GitHub
1452def training_step_seconds(params: float, tokens: int, flops_per_second: float = PEAK_FLOPS) -> float:
1453    """Lower bound on one GPU's arithmetic for a training step: 6 FLOPs per parameter per token."""
1454    return 6 * params * tokens / flops_per_second

Lower bound on one GPU's arithmetic for a training step: 6 FLOPs per parameter per token.

LLAMA3_70B_SHAPE = {'layers': 80, 'kv_heads': 8, 'head_dim': 128}
def max_cache_tokens( gpu_bytes: float, params: float, weight_bits: int, kv_bits: int, layers: int, kv_heads: int, head_dim: int) -> int: on GitHub
1465def max_cache_tokens(gpu_bytes: float, params: float, weight_bits: int, kv_bits: int,
1466                     layers: int, kv_heads: int, head_dim: int) -> int:
1467    """Tokens of KV cache that fit beside the weights (0 if the weights alone don't fit).
1468
1469    The formulas are primer.ml.inference's; this lesson varies the number format.
1470    """
1471    free = gpu_bytes - weight_bytes(params, weight_bits)
1472    if free <= 0:
1473        return 0
1474    return int(free // kv_cache_bytes_per_token(layers, kv_heads, head_dim, kv_bits))

Tokens of KV cache that fit beside the weights (0 if the weights alone don't fit).

The formulas are primer.ml.inference's; this lesson varies the number format.

def figures() -> dict: on GitHub
1489def figures() -> dict:
1490    """Plot this lesson's data. matplotlib is imported here, and only here,
1491    so the lesson itself needs nothing beyond NumPy."""
1492    import matplotlib
1493
1494    matplotlib.use("Agg")
1495    import matplotlib.pyplot as plt
1496
1497    BLUE, RED, GREEN, ORANGE, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1498    figs = {}
1499
1500    # --- 1. Parallel rounds: speed-up until the independent work runs out ----
1501    cores = 2 ** np.arange(0, 15)
1502    fig, ax = plt.subplots(figsize=(6, 3.4))
1503    ax.plot(cores, [parallel_rounds(64, 64, 64, int(c)) for c in cores], "o-", color=BLUE)
1504    ax.set_xscale("log", base=2)
1505    ax.set_yscale("log")
1506    ax.axvline(64 * 64, color=MUTED, ls="--")
1507    ax.text(64 * 64 / 1.3, 1e4, "cores = cells\n(4,096)", color="#4b5563", ha="right")
1508    ax.set(xlabel="cores working at once", ylabel="rounds to finish",
1509           title="A 64 × 64 × 64 multiply: faster until every cell has a core")
1510    figs["parallel_rounds"] = fig
1511
1512    # --- 2. The memory hierarchy: capacity up, bandwidth down --------------
1513    names = [lv.name for lv in MEMORY_HIERARCHY]
1514    y = np.arange(len(names))[::-1]
1515    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4), sharey=True)
1516    a1.barh(y, [lv.capacity_bytes for lv in MEMORY_HIERARCHY], color=BLUE)
1517    a2.barh(y, [lv.bandwidth_bytes_per_s for lv in MEMORY_HIERARCHY], color=ORANGE)
1518    a1.set_yticks(y, names)
1519    a1.set(xscale="log", xlabel="capacity (bytes)", title="Further out holds more...")
1520    a2.set(xscale="log", xlabel="bandwidth (bytes per second)", title="...and moves it more slowly")
1521    for a in (a1, a2):
1522        a.grid(axis="y", visible=False)
1523    fig.text(0.5, -0.02, "Illustrative orders of magnitude for one datacenter GPU, not a product specification",
1524             ha="center", color="#4b5563", fontsize=8)
1525    fig.tight_layout()
1526    figs["hierarchy"] = fig
1527
1528    # --- 3. Tiling: traffic falls, intensity rises --------------------------
1529    n_sim = 32
1530    tiles = [1, 2, 4, 8, 16, 32]
1531    measured = [tiled_matmul(np.ones((n_sim, n_sim)), np.ones((n_sim, n_sim)), t)[1].reads for t in tiles]
1532    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1533    a1.loglog(tiles, [matmul_reads(n_sim, t) for t in tiles], color=MUTED, base=2, label="formula 2n³/T")
1534    a1.loglog(tiles, measured, "o", color=BLUE, base=2, label="counted by tiled_matmul")
1535    a1.set(xlabel="tile width T", ylabel="numbers read from HBM", title=f"Reads for two {n_sim} × {n_sim} matrices")
1536    a1.legend(frameon=False)
1537    big = 4096
1538    ts = 2 ** np.arange(0, 12)
1539    for b, color, label in ((2, BLUE, "16-bit"), (1, GREEN, "8-bit")):
1540        a2.plot(ts, [matmul_intensity(big, int(t), b) for t in ts], "o-", color=color, label=label)
1541    a2.set_xscale("log", base=2)
1542    a2.set_yscale("log")
1543    ridge = PEAK_FLOPS / HBM_BANDWIDTH
1544    a2.axhline(ridge, color=RED, ls="--")
1545    a2.text(1.2, ridge * 1.25, f"break-even ≈ {ridge:.0f} FLOPs/byte", color=RED)
1546    a2.set(xlabel="tile width T", ylabel="FLOPs per byte moved", title=f"Intensity of a {big} × {big} multiply")
1547    a2.legend(frameon=False, loc="lower right")
1548    fig.tight_layout()
1549    figs["tiling"] = fig
1550
1551    # --- 4. Bit layouts, drawn to scale -------------------------------------
1552    rows = [(f.name, [("sign", 1), ("exponent", f.exponent_bits), ("mantissa", f.mantissa_bits)]) for f in FLOAT_FORMATS]
1553    rows += [("int8", [("integer", 8)]), ("int4", [("integer", 4)])]
1554    colors = {"sign": MUTED, "exponent": ORANGE, "mantissa": BLUE, "integer": GREEN}
1555    fig, ax = plt.subplots(figsize=(9, 3.6))
1556    for r, (name, parts) in enumerate(rows):
1557        yy, x0 = len(rows) - 1 - r, 0
1558        for part, width in parts:
1559            ax.barh(yy, width, left=x0, height=0.7, color=colors[part], edgecolor="white", linewidth=0.5)
1560            if width >= 2:
1561                ax.text(x0 + width / 2, yy, str(width), ha="center", va="center", color="white", fontsize=8)
1562            x0 += width
1563        ax.text(x0 + 0.4, yy, f"{x0} bits", va="center", fontsize=8, color="#4b5563")
1564    ax.set_yticks(range(len(rows)), [name for name, _ in rows][::-1])
1565    ax.set_xlim(0, 36)
1566    ax.set_xlabel("bits")
1567    ax.grid(False)
1568    handles = [plt.Rectangle((0, 0), 1, 1, color=c) for c in colors.values()]
1569    ax.legend(handles, list(colors), frameon=False, loc="lower right")
1570    ax.set_title("Where each format spends its bits")
1571    figs["format_layouts"] = fig
1572
1573    # --- 5. Relative precision across magnitudes ----------------------------
1574    x = np.logspace(-10, 6, 1200)
1575    fig, ax = plt.subplots(figsize=(8.5, 4))
1576    top = 10  # a line that would climb past the chart ends at its top instead of running into the title
1577
1578    def within(y):
1579        return np.where(y <= top, y, np.nan)
1580
1581    for fmt, color in zip(FLOAT_FORMATS, (BLUE, GREEN, ORANGE, RED, "#7c3aed")):
1582        ax.loglog(x, within(_relative_spacing(x, fmt)), color=color, label=fmt.name)
1583    for bits, ls in ((8, "--"), (4, ":")):
1584        step = 1 / (2 ** (bits - 1) - 1)  # scale chosen so 1.0 lands on the top code
1585        ax.loglog(x, within(np.where(x <= 1, step / x, np.nan)), color="#4b5563", ls=ls, label=f"int{bits}, scale 1/{2 ** (bits - 1) - 1}")
1586    ax.set_ylim(1e-8, top)
1587    ax.set(xlabel="size of the number stored", ylabel="gap to the next value ÷ the number",
1588           title="Relative precision: flat for floats, rising for integers")
1589    ax.legend(frameon=False, fontsize=8, loc="center left", bbox_to_anchor=(1.01, 0.5))
1590    fig.tight_layout()
1591    figs["precision"] = fig
1592
1593    # --- 6. Communication vs. computation -----------------------------------
1594    gpus = 2 ** np.arange(1, 11)
1595    grad_bytes, params, tokens = 14e9, 7e9, 16_384
1596    fig, ax = plt.subplots(figsize=(6.5, 3.6))
1597    ax.semilogx(gpus, [all_reduce_seconds(grad_bytes, int(g), BETWEEN_MACHINES_LINK) for g in gpus], "o-", color=RED,
1598                base=2, label="all-reduce over the network (50 GB/s)")
1599    ax.semilogx(gpus, [all_reduce_seconds(grad_bytes, int(g), IN_MACHINE_LINK) for g in gpus], "o-", color=BLUE,
1600                base=2, label="all-reduce over in-machine links (500 GB/s)")
1601    ax.axhline(training_step_seconds(params, tokens), color=MUTED, lw=3, label="arithmetic per step (16k tokens)")
1602    ax.set_ylim(0, 0.8)
1603    ax.set(xlabel="GPUs in the all-reduce", ylabel="seconds per training step",
1604           title="7B parameters, 16-bit gradients: talking vs. computing")
1605    ax.legend(frameon=False, fontsize=8, loc="center right")
1606    figs["communication"] = fig
1607
1608    # --- 7. Will it fit? ----------------------------------------------------
1609    gpu, params = 80e9, 70e9
1610    choices = [("16-bit", 16), ("fp8", 8), ("4-bit", 4)]
1611    fig, ax = plt.subplots(figsize=(6, 3.6))
1612    for i, (label, bits) in enumerate(choices):
1613        w = weight_bytes(params, bits) / 1e9
1614        free = max(gpu / 1e9 - w, 0)
1615        ax.bar(i, w, color=BLUE, label="weights" if i == 0 else None)
1616        ax.bar(i, free, bottom=w, color="#93c5fd", label="room for KV cache" if i == 0 else None)
1617        tokens_fit = max_cache_tokens(gpu, params, bits, 16, **LLAMA3_70B_SHAPE)
1618        note = "does not fit" if tokens_fit == 0 else f"{tokens_fit:,} tokens"
1619        ax.text(i, max(w, gpu / 1e9) + 3, note, ha="center", fontsize=9)
1620    ax.axhline(gpu / 1e9, color=RED, ls="--", label="80 GB of GPU memory")
1621    ax.set_xticks(range(len(choices)), [f"{label} weights" for label, _ in choices])
1622    ax.set_ylim(0, 185)
1623    ax.set(ylabel="GB", title="A 70B model on one GPU (16-bit KV cache)")
1624    ax.legend(frameon=False, loc="upper right")
1625    figs["serving_fit"] = fig
1626
1627    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1635def demo() -> None:
1636    banner("1. The workload: a matrix multiply, counted by hand")
1637    C, muls, adds = counted_matmul([[1, 2, 3], [4, 5, 6]], [[7, 8], [9, 10], [11, 12]])
1638    say(f"[[1,2,3],[4,5,6]] times [[7,8],[9,10],[11,12]] = {C}.")
1639    say(
1640        f"""
1641        That took {muls} multiplies and {adds} adds: 2·m·n·k = {matmul_flops(2, 2, 3)} FLOPs.
1642        No output cell needed another cell's answer, so they could all be
1643        computed at once. A 4096-cube multiply is {matmul_flops(4096, 4096, 4096):.3e} FLOPs.
1644        """
1645    )
1646    table(["cores", "rounds for a 64 × 64 × 64 multiply"], [(c, parallel_rounds(64, 64, 64, c)) for c in (1, 64, 4096, 8192)])
1647    takeaway("A GPU wins by doing thousands of independent multiply-adds at once; it needs that much independent work.")
1648
1649    banner("2. The memory hierarchy")
1650    table(
1651        ["level", "holds", "moves per second", "ms to stream 1 GB", "picture"],
1652        [(lv.name, f"{lv.capacity_bytes:.0e} B", f"{lv.bandwidth_bytes_per_s:.2e} B", transfer_seconds(1e9, lv.name) * 1e3, lv.picture)
1653         for lv in MEMORY_HIERARCHY],
1654        floatfmt=".2f",
1655    )
1656    say(
1657        """
1658        Illustrative round numbers. Reading 16 GB of weights from HBM takes
1659        4.78 ms; from host memory it would take 320 ms.
1660        """
1661    )
1662    takeaway("Each step away from the arithmetic is bigger and slower; full speed means living in HBM or closer.")
1663
1664    banner("3. Tiling: same FLOPs, far fewer bytes")
1665    rows = []
1666    for t in (1, 2, 4):
1667        _, traffic = tiled_matmul(np.ones((4, 4)), np.ones((4, 4)), t)
1668        rows.append((f"{t} × {t}", traffic.reads, traffic.writes, traffic.flops, traffic.flops / (traffic.reads + traffic.writes)))
1669    table(["tile", "reads", "writes", "FLOPs", "FLOPs per number moved"], rows, floatfmt=".2f")
1670    say(
1671        f"""
1672        For 4096 × 4096 matrices with 128-wide tiles: {matmul_intensity(4096, 128, 2):.1f} FLOPs per byte
1673        in 16-bit, {matmul_intensity(4096, 128, 1):.1f} in 8-bit. The chip breaks even at
1674        {PEAK_FLOPS / HBM_BANDWIDTH:.0f}. FlashAttention, kernel fusion and batching all
1675        raise this ratio.
1676        """
1677    )
1678    takeaway("Speed is usually set by bytes moved, not FLOPs; reuse every byte you fetch as often as you can.")
1679
1680    banner("4. Number formats, built from scratch")
1681    say(f"-6.5 in bf16 is stored as {bit_string(-6.5, BF16)} (sign, exponent, mantissa).")
1682    table(
1683        ["format", "bits", "largest", "smallest normal", "gap above 1", "0.1 is stored as", "multiplier cells"],
1684        [(f.name, f.bits, f"{f.max_value:.4g}", f"{f.min_normal:.3g}", f"{f.epsilon:.3g}", f"{round_to(0.1, f):.8g}", multiplier_cells(f))
1685         for f in FLOAT_FORMATS],
1686    )
1687    say(
1688        f"""
1689        Range: 1e-8 in fp16 becomes {round_to(1e-8, FP16)}; in bf16 it is {round_to(1e-8, BF16):.4g}.
1690        Precision: 1 + 0.001 in bf16 is {round_to(1.001, BF16)}; in fp32 it is {round_to(1.001, FP32):.7g}.
1691        Overflow: 70,000 in fp16 is {round_to(70_000.0, FP16)}; 500 in fp8 E4M3 is {round_to(500.0, FP8_E4M3)}.
1692        """
1693    )
1694    takeaway("Exponent bits buy range, mantissa bits buy precision; fewer bits buy speed three ways.")
1695
1696    banner("5. Many GPUs: ring all-reduce")
1697    results, sent = ring_all_reduce([np.arange(8.0) * (g + 1) for g in range(4)])
1698    say(f"Four GPUs, 8 numbers each. After the ring, every GPU holds {results[0].tolist()}; each sent {sent[0]} numbers.")
1699    table(
1700        ["link", "all-reduce of 14 GB on 8 GPUs (ms)", "arithmetic per step (ms)"],
1701        [("in one machine", all_reduce_seconds(14e9, 8, IN_MACHINE_LINK) * 1e3, training_step_seconds(7e9, 16_384) * 1e3),
1702         ("between machines", all_reduce_seconds(14e9, 8, BETWEEN_MACHINES_LINK) * 1e3, training_step_seconds(7e9, 16_384) * 1e3)],
1703        floatfmt=".0f",
1704    )
1705    takeaway("Put chatty parallelism on fast links; send the once-per-step traffic over the network.")
1706
1707    banner("6. Will a 70B model fit on one 80 GB GPU?")
1708    table(
1709        ["weights", "weight GB", "tokens of 16-bit cache", "tokens of 8-bit cache"],
1710        [(f"{b}-bit", weight_bytes(70e9, b) / 1e9, max_cache_tokens(80e9, 70e9, b, 16, **LLAMA3_70B_SHAPE),
1711          max_cache_tokens(80e9, 70e9, b, 8, **LLAMA3_70B_SHAPE)) for b in (16, 8, 4)],
1712        floatfmt=".0f",
1713    )
1714    takeaway("The number format decides whether a model fits, how many users share a GPU, and how fast each token comes.")