The lesson in one minute
What you'll be able to explain
- 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.
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
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:
- 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.
- 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.
- 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?
Chapter 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
Symbols
| Symbol | Meaning here | Shape |
|---|---|---|
| the left matrix: 2 rows of 3 numbers | 2 × 3 | |
| the right matrix: 3 rows of 2 numbers | 3 × 2 | |
| the matrix multiply: the cell in row , column is row of dotted with column of | 2 × 2 | |
| 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
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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| rows of the left matrix (and of the answer) | 2 | |
| the shared inner size: columns of the left, rows of the right; the length of each dot product | 3 | |
| columns of the right matrix (and of the answer) | 2 | |
| how many cells the answer has | 4 | |
| 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
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.
Figure 2 · Diagram
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
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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| independent cells to compute | 64 × 64 = 4,096 | |
| cores | identical cores working at once | 1 to 8,192 |
| ceiling: round up to a whole number (half a cell of work still needs a round) | ||
| 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
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)
Figure 1 · Drawn from the lesson's code
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
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.
Chapter 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
Symbols
| Symbol | Meaning here | Units |
|---|---|---|
| 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
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
Figure 4 · Diagram
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"]
Figure 3 · Drawn from the lesson's code
Capacity grows about fifty-million-fold from registers to the network while bandwidth falls ten-thousand-fold from registers to disk
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.
Chapter 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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| both matrices are | 4, then 4,096 | |
| tile width: fast memory works on blocks | 1, 2, 4, then 128 | |
| the FLOPs of the whole multiply (section 1, with ) | 128 | |
| reads | numbers fetched from slow memory | 128, 64, 32 |
| numbers written back: each answer cell once | 16 | |
| bytes per number: 2 at 16-bit, 1 at 8-bit | 2 | |
| arithmetic intensity: FLOPs per byte moved | FLOPs/byte | |
| "roughly", once is much bigger than 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
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
Figure 6 · Diagram
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
Figure 5 · Drawn from the lesson's code
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
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.
Chapter 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).
- Sign: negative, so the sign bit is 1.
- Power of two: the largest power of two not above 6.5 is 4 = 2², so 6.5 = 1.625 × 2².
- Exponent: stored with a bias of 127 added, so that negative powers need no sign of their own: 2 + 127 = 129 = 10000001 in binary.
- 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.
- 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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the sign bit: 0 positive, 1 negative | 1 | |
| −1 multiplied by itself times: +1 when , −1 when | −1 | |
| how many exponent bits the format has | 8 | |
| the stored exponent, read as an ordinary whole number | 129 | |
| bias | the offset subtracted from , so stored values 1…254 stand for powers −126…127 | 127 |
| how many mantissa bits the format has | 7 | |
| the stored mantissa, read as a whole number from 0 to | 1010000 = 80 | |
| the significand: the hidden leading 1 plus the stored fraction, between 1 and 2 | 1 + 80/128 = 1.625 | |
| 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
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.
Figure 9 · Diagram
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
Figure 7 · Drawn from the lesson's code
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
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.
Figure 8 · Drawn from the lesson's code
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
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:
- 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. - Intensity. The same tile does twice the FLOPs per byte (section 3: 63 becomes 126), pushing more work past the break-even point.
- 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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| stored mantissa bits | 23, 10, 7, 3 | |
| 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
# 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.
Chapter 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.
- 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.
- 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.
Figure 11 · Diagram
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
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| GPUs taking part | 8 | |
| bytes of gradient each GPU holds | 7 × 10⁹ parameters × 2 bytes = 14 GB | |
| each GPU's link bandwidth | 500 GB/s in one machine, 50 GB/s between machines (illustrative) | |
| the fraction of its gradient each GPU sends in a ring all-reduce: just under 2 | 1.75 | |
| parameters in the model | 7 × 10⁹ | |
| tokens each GPU processes per step | 16,384 | |
| FLOPs per parameter per token in training: 2 forward, 4 backward | ||
| 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
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
Figure 10 · Drawn from the lesson's code
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
Figure 12 · Diagram
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
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.
Chapter 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
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| parameters | 7 × 10¹⁰ | |
| bytes per weight: bits ÷ 8 | 2, 1 or 0.5 | |
| tokens held in the KV cache, summed over every request being served | the unknown | |
| one key and one value per token | ||
| layers, each with its own cache | 80 | |
| key/value heads | 8 | |
| dimensions per head | 128 | |
| bytes per cached number | 2 or 1 | |
| the GPU's memory | 80 × 10⁹ bytes | |
| "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
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)
Figure 14 · Diagram
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"]
Figure 13 · Drawn from the lesson's code
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
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.
Test yourself
8 questions
Answer each one out loud or on paper before you open it. If you can explain it, you know it.
Question 1Why are GPUs, rather than CPUs, used to train and run neural networks?Think it through, then reveal
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.
Question 2How many FLOPs does multiplying a 1,000 × 2,000 matrix by a 2,000 × 500 matrix take, and why that formula?Think it through, then reveal
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.
Question 3The chip can do 10¹⁵ FLOPs per second but a model runs far slower. What is usually the bottleneck, and how do you tell?Think it through, then reveal
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.
Question 4How does tiling a matrix multiply reduce memory traffic, and what limits it?Think it through, then reveal
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.
Question 5What is the difference between bf16 and fp16, and why do many training runs prefer bf16?Think it through, then reveal
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.
Question 6Why does halving the bits per number roughly double throughput?Think it through, then reveal
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.
Question 7Why is tensor parallelism usually kept inside one machine while data parallelism spans many?Think it through, then reveal
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.
Question 8Will a 70-billion-parameter model serve from one 80 GB GPU?Think it through, then reveal
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.
Primary sources
The papers behind this lesson
Introduced the roofline: judge a kernel by its arithmetic intensity against the machine's balance of compute and bandwidth.
The paper ↗Applied tiling to attention so the score matrix never reaches HBM, showing that counting memory traffic rather than FLOPs is what makes attention fast.
Read the annotated companion →The paper ↗Showed that networks train in 16-bit floats with fp32 master weights, fp32 accumulation and loss scaling.
The paper ↗Showed that bf16, with fp32's range, trains a wide range of models to fp32 quality without loss scaling.
The paper ↗Proposed the E4M3 and E5M2 8-bit formats and showed that training and inference hold up in them.
The paper ↗Split each transformer layer's matrix multiplies across the GPUs of one machine, the tensor parallelism of section 5.
Read the annotated companion →The paper ↗Brought the bandwidth-optimal ring all-reduce to deep learning training.
The paper ↗Researcher's shelf
Further reading
- Horace He, Making Deep Learning Go Brrrr From First Principles: https://horace.io/brrr_intro.html
- How to Scale Your Model (a book on TPUs, GPUs and parallelism for transformers): https://jax-ml.github.io/scaling-book/
- Williams, Waterman and Patterson, Roofline (2009): https://doi.org/10.1145/1498765.1498785
- Dao et al., FlashAttention (2022): https://arxiv.org/abs/2205.14135
- Micikevicius et al., Mixed Precision Training (2017): https://arxiv.org/abs/1710.03740
- Micikevicius et al., FP8 Formats for Deep Learning (2022): https://arxiv.org/abs/2209.05433
- PyTorch automatic mixed precision: https://pytorch.org/docs/stable/amp.html
- NVIDIA, CUDA C++ Programming Guide (how one GPU family organises cores and memory): https://docs.nvidia.com/cuda/cuda-programming-guide/index.html
About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.