How much memory and time does a model need to write one token?
A short test to skip this level
Solve these 6 questions on your own. Answer all of them correctly and the level counts as cleared with three stars, and every part of the page opens. Showing an answer doesn’t count.
Number A 24 GiB GPU, 16 GiB of weights, the rest for the KV cache. 4 K/V heads, dₖ = 128, 32 layers, 2 bytes per number. How many tokens of cache fit?
Number One block: model width 64, FFN width 256, and the new token attends to 128 tokens. How many FLOPs does this token cost in the block? (Count the matrix products only.)
Number 7 billion weights of 2 bytes each. The GPU reads 10¹² bytes per second. One user: at most how many tokens per second? (1 decimal)
Number A 7-billion-weight model is stored in 1 byte per weight. The GPU reads 10¹² bytes and does 10¹⁴ FLOPs per second, and each user’s token costs 1.4 × 10¹⁰ FLOPs. For which batch size B does the arithmetic take as long as the read?
Number A 7-billion-weight model, 2 bytes per weight, runs on a GPU that reads 10¹² bytes and does 10¹⁴ FLOPs per second. The prompt has 50 tokens. How many ms does the prefill pass take?
CodeWrite serve: for one decode step with B users, return the step time in ms, the tokens per second for each user, and the tokens per second in total. The step reads the weights and every user’s cache, and does B × 2N FLOPs (the weights only, not attention’s part). bytes_per applies to the weights and the cache. Reading and arithmetic overlap, so the step takes the longer one.
Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs
Warm-up2 questions from earlier levels
A quick review before you start. Optional. Nothing here locks the level.
Levels 16 to 19 built a model that gives correct answers. This level asks what it costs to run. When a model writes, two things
limit it. One is the memory it needs on the GPU. The other is the time each new token takes. Both come from a
few numbers about the model, with no training at all. You can check every number below by hand.
1. The KV cache: what the model keeps while it writes
When a model writes, each new token attends to every earlier token, so it needs their K and V.
Those never change once computed, because a token never sees anything after it. So the model keeps them in GPU memory.
For each new token, it computes q, k and v for that one token only:
token given to the model (counting from 0)
computed now
read from the KV cache
token 3
q, k, v of token 3
k, v of tokens 0–2
token 4
q, k, v of token 4
k, v of tokens 0–3
This is the same cache as in level 18, section 7. Now count its size. In every layer and every attention head, each
token keeps one k vector and one v vector of dk numbers each. Models store each number in 16 bits, which is
2 bytes. Memory is counted in KiB (1 KiB = 1,024 bytes) and GiB (1 GiB = 1,024 × 1,024 KiB):
bytes per token=2×heads×dk×layers×2
The first 2 is for k and v; the last 2 is the bytes per number. A tiny case: 1 head, dk=4, 2 layers gives
2 × 1 × 4 × 2 × 2 = 32 bytes per token.
Number A model has 32 layers and 32 attention heads with dₖ = 128, and every head keeps its own K and V. Each number takes 2 bytes. How many KiB of KV cache does one token need? (1 KiB = 1,024 bytes)
I got stuck here Why does the cache grow with the conversation, and not just with the model?
Every token written so far keeps its own k and v, in every layer. The 100th token of an answer attends to
the 99 tokens before it (and the whole prompt). So all of them stay in memory until the conversation ends.
If the conversation is twice as long, the cache is twice as big. The weights, in contrast, are the same size
whether the conversation is 10 tokens or 10,000.
2. Grouped-query attention: share K and V
That cache grows with every token, every layer and every K/V head. It often limits how many users one GPU
can serve. Grouped-query attention makes it smaller: several query heads share one K/V head.
🔒 Answer the question above to unlock
Grouped-query attention: query heads share K/V heads, so the cache shrinks
Pick how many K/V heads there are. Point at or tap a K/V head to see which query heads read it.
8 K/V heads, one per query head
K/V cache per token8,192 bytes
A small model, to keep the picture readable: 8 query heads, dₖ = 64, 4 layers. The questions use 32 query heads and 32 layers, but the sharing works the same way.
K/V heads
K and V stored per token = 2 × 8 heads × 64 × 4 layers × 2 bytes = 8,192 bytes(100.0% of the 8-head size)
the group you picked: one K/V head and the query heads that read itother groups
With fewer K/V heads, several query heads use the same keys and values. Each query head still has its own q, so each
one can still look for something different. But they all search the same keys.
Number A model has 32 query heads and 8 K/V heads, numbered from 0; consecutive query heads share a K/V head. Which K/V head does query head 13 use?
Predict firstGuess first: going from 32 K/V heads to 1 (all 32 query heads share one K and one V) cuts the KV cache by 32×. What happens to quality?
🔒 Answer the question above to unlock
Try it
In the lab, go from 8 K/V heads to 1. The memory per token drops to 12.5%. Then try 4 and 2. Which choice would you
make if the cache, not the weights, filled your GPU? What would the query heads lose?
Number A GPU has 24 GiB of memory, and the model’s weights take 16 GiB of it. The rest holds the KV cache. The model uses grouped-query attention with 4 K/V heads, dₖ = 128 and 32 layers, and stores each number in 2 bytes. How many tokens of KV cache fit? (1 KiB = 1,024 bytes, 1 GiB = 1,024 × 1,024 KiB)
All of this is the memory for writing. Training needs much more memory for each weight, and it keeps values from the
forward pass for the backward pass. Side trip U6 counts that memory.
3. FLOPs: how much arithmetic one token costs
Work is counted in FLOPs: floating-point operations, each one multiplication or one addition. A matrix product
(1, a) @ (a, b) gives b numbers. Each of them needs a multiplications and about a additions.
So the product costs 2 × a × b FLOPs. For (1, 3) @ (3, 4) that is 24.
A token passes through every weight matrix of the model once, and each weight does one multiply and one add.
So a model with N weights costs about 2N FLOPs per token in the forward pass. Looking up a token’s row in the
embedding table is not a multiplication, so strictly N counts only the weights used in matrix products.
The numbers in this level are large, so they are written as powers of ten. To divide powers of ten, subtract the
exponents: 10¹⁴ / 10¹⁰ = 10⁴. The math page shows more examples.
🔒 Answer the question above to unlock
Number A model has 7 billion weights. A GPU does 10¹⁴ FLOPs per second (100 trillion). Using the 2N rule for the forward pass, at most how many tokens per second can it produce, counting only the arithmetic?
🔒 Answer the question above to unlock
One part is missing from 2N: attention itself. Inside one block, a new token with vector size d (the model width)
does the steps below. It attends to L tokens, and all heads together use d numbers per token.
step
product
FLOPs per token
q, k, v and the output: four (d,d) matrices (one K/V head per query head)
(1,d)@(d,d), four times
2×4d2
the FFN: (d,dff), then (dff,d)
(1,d)@(d,dff), then (1,dff)@(dff,d)
2×2ddff
scores: q against the k of L tokens
(1,d)@(d,L)
?
weighted sum of L value vectors
(1,L)@(L,d)
?
The first two rows are the 2N rule for the block’s weights. With grouped-query attention, the k and v matrices are
smaller than (d,d), and 2N still counts their weights correctly. The last two rows use no weights at all. Use the
2 × a × b rule on their shapes.
Number One block has model width d = 64 and FFN width 256. A new token attends to L = 128 tokens. Its attention needs two products: the scores, (1, 64) @ (64, 128), and the weighted sum, (1, 128) @ (128, 64). How many FLOPs does this token cost in the whole block, weights and attention together?
🔒 Answer the question above to unlock
Each attention product costs 2Ld FLOPs, so the attention part is 4Ld. It grows with L, the number of tokens the new
token attends to. The weights part stays the same.
Number One block has model width 64 and FFN width 256, so its weights part is 98,304 FLOPs per token. Its attention part is 4 × L × 64 FLOPs. At what context length L are the two parts equal?
Go deeper Training costs about 6N per token
Training runs the forward pass (2N) and the backward pass, which costs about twice as much. For each weight, the
backward pass computes the gradient for that weight and the gradient for its input (the output of the layer before). So training costs about
6N FLOPs per token. Level 21’s reference GPT has N = 102,415 weights: about 0.2 million FLOPs per token forward.
Its two embedding tables hold about 2% of the weights, so counting them barely changes the estimate.
One epoch is 8,500 problems of 10 tokens each. So one epoch costs about 6 × 102,415 × 85,000 ≈ 5.2 × 10¹⁰ FLOPs.
A laptop can do that arithmetic in seconds. The real run takes longer because of Python and because the matrices
are small. Level 21’s context is only 10 tokens, so the attention part is tiny there.
4. Decode: one token per step
A GPU keeps the weights in its memory, and the arithmetic units must read them from there. How fast it can read is the
memory bandwidth, in bytes per second. Bandwidth is given in decimal units, so in the rest of this level 1 GB = 10⁹ bytes.
(A GiB is 1,024³ bytes, about 1.07 GB.)
Writing the answer one token at a time is called decode. Each decode step makes one new token.
🔒 Answer the question above to unlock
Predict firstBefore you compute it: a GPU computes 10¹⁴ FLOPs per second and reads 10¹² bytes per second from its memory. It writes one user’s answer with a 7-billion-weight model. What limits the speed?
🔒 Answer the question above to unlock
Every weight is read once for each decode step. The next token needs this one as its input. So it can’t start before
this one is done, and the weights are read again.
For one user, that gives a simple upper limit:
tokens per second≤weight bytesbandwidth
Number A model has 7 billion weights of 2 bytes each. The GPU reads 10¹² bytes per second from its memory and must read every weight once for each new token. For one user, at most how many tokens per second can it write? (1 decimal)
🔒 Answer the question above to unlock
A server doesn’t write for one user at a time. In one step it reads the weights once. It uses them for the next
token of B users together: B tokens for the cost of one read. The arithmetic still grows with B. So for a large enough
B, the arithmetic, not the read, is the limit. The step takes the longer of the two times.
Number A model has 7 billion weights of 2 bytes each. One decode step reads the 1.4 × 10¹⁰ bytes of weights at 10¹² bytes per second. Each user’s token costs 1.4 × 10¹⁰ FLOPs at 10¹⁴ FLOPs per second. One step serves B users at once, with one read of the weights. For which B does the arithmetic take as long as the read?
🔒 Answer the question above to unlock
The B where the two times are equal is called the crossover. Below it, reading the weights is the limit; above it,
the arithmetic is.
Each step also reads the KV cache of every user. Take 8 users with 4,000 tokens each, and 128 KiB per token.
Their caches hold 4,096,000 KiB. Times 1,024, that is about 4.2 × 10⁹ bytes, or 4.2 GB. So one step reads 14 + 4.2 ≈ 18.2 GB and takes 18.2 ms instead of
14 ms. Each user now gets about 55 tokens per second instead of 71.
The caches must also fit in the GPU’s memory, next to the weights. Section 2 counted memory in GiB. Here, as for
bandwidth, 1 GB = 10⁹ bytes.
Number A GPU has 80 GB of memory. It holds the 14 GB of weights of the 7-billion-weight model, and the KV cache of every user. Each user has 4,000 tokens of cache, at 128 KiB per token. At most how many users fit? (1 GB = 10⁹ bytes, 1 KiB = 1,024 bytes)
🔒 Answer the question above to unlock
One decode step for B users: reading time and arithmetic time
Move the slider to add users. Then give each user a longer conversation.
tokens per user
B = 8: read 14.00 ms, arithmetic 1.12 ms, so a step takes 14.00 ms. 71.4 tokens/s per user,571 in total. Memory: 14 GB of weights + 0.0 GB of cache= 14.0 GB of 80.
your Bcrossover: the two times are equal at B = 100
I got stuck here If extra users cost almost nothing, why not serve 1,000 at once?
There are two reasons. Above B = 100, each extra user adds arithmetic to every step, so every user gets tokens more
slowly. With 256 users a step takes 256 × 0.14 ms ≈ 36 ms instead of 14 ms. Also, every user brings a KV cache,
which must fit in memory next to the weights. With 4,000 tokens per user, only 125 users fit on an 80 GB GPU. That is
why the KV cache, not the arithmetic, often decides how many users one GPU serves.
The weights are read once for all B users. Each user’s cache, in contrast, is read for that user alone, so its reads
can’t be shared across users. With long conversations, extra users no longer cost almost nothing. In the lab, give
each user 4,000 tokens: the two lines never meet. Fewer K/V heads make the cache smaller, which saves time as well as
memory.
Another way to read less is to store each weight in fewer bytes.
Number Store each of the 7 billion weights in 1 byte instead of 2. The GPU still reads 10¹² bytes per second, and each user’s token still costs 1.4 × 10¹⁰ FLOPs at 10¹⁴ FLOPs per second. For which B does the arithmetic now take as long as the read?
5. Prefill: reading the prompt in one pass
Before decode starts, the model must read the prompt. All prompt tokens are known at once. So they pass through the
model in one forward pass, the same way a whole sequence does in training (level 17, section 4).
This pass is called prefill. It also fills the KV cache with the k and v of every prompt token.
In prefill, the weights are read once for all L prompt tokens. That is the same idea as a batch of L users.
🔒 Answer the question above to unlock
Number A model has 7 billion weights of 2 bytes each. The GPU reads 10¹² bytes per second and does 10¹⁴ FLOPs per second. It reads a prompt of 2,000 tokens in one pass. The weights are read once (14 ms), and each token costs 1.4 × 10¹⁰ FLOPs. How many ms does the pass take?
🔒 Answer the question above to unlock
Number The same model and GPU: 7 billion weights of 2 bytes, 10¹² bytes per second, 10¹⁴ FLOPs per second. Now the prompt has only 50 tokens. How many ms does the prefill pass take?
🔒 Answer the question above to unlock
So a chat answer has two parts with different limits. The first new token waits for prefill. For a prompt longer than
about 100 tokens (the B = 100 of section 4), prefill is limited by arithmetic. A shorter prompt still waits for the
weight read. Every token after the first comes from decode, which is limited by memory.
Now put the whole decode step into one function. It is the estimate you can make for any model before you run it.
CodeWrite serve for one decode step. B users each have a KV cache of the given number of tokens. The step reads the weights once and every user’s cache. It also does 2N FLOPs for each user’s token. bytes_per is the bytes of every stored number: the weights and the KV cache. The step takes the longer of the read and the arithmetic. Return the step time in ms, the tokens per second for each user, and the tokens per second in total.
Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs
Level 21 puts the parts together: you write and train your own GPT. Its context is only 10 tokens, so its
attention part and its KV cache are tiny.
Recap
a summary for when you finish the level
The key formulas and common mistakes appear here once you clear the level.
You can now
Size the KV cache of a model, and see how grouped-query attention makes it smaller.
Count the FLOPs of one token: 2N for the weights, plus an attention part that grows with the context.
Find the time of one decode step for B users, and distinguish prefill (limited by arithmetic, for a long prompt) from decode (limited by memory).
Keep in mind
KV bytes per token=2×K/V heads×dk×layers×2
FLOPs per token≈2N; one block: 2(4d2+2ddff)+4Ld
read=bandwidthweight bytes+B×cache bytes
step=max(read,FLOPs per secondB×2N)
Prefill reads the whole prompt in one pass. For a prompt longer than about 100 tokens, it waits for arithmetic. Decode writes one token per step and waits for memory.
Common mistakes
Forgetting the 2 for K and V, or the 2 bytes per number, when you size the KV cache.
Counting one FLOP per weight instead of two (a multiply and an add).