# 20. Inference cost: memory and speed

> How much memory and time does a model need to write one token?

LLM by Hand · Theory · runs in your browser · interactive page: https://llm.liko.page/learn/inference-cost/

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 $d_k$ 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):

$$
\text{bytes per token} = 2 \times \text{heads} \times d_k \times \text{layers} \times 2
$$

The first 2 is for k and v; the last 2 is the bytes per number. A tiny case: 1 head, $d_k = 4$, 2 layers gives
2 × 1 × 4 × 2 × 2 = 32 bytes per token.

**Question.** 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)

*Answer it on the page to check your work.*

**If you are stuck: 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.

*[Interactive lab: Gqa — open the page to use it]*

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.

**Question.** 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?

*Answer it on the page to check your work.*

**Predict.** Guess 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?

A. Nothing changes: K and V hold the same information either way
B. It drops a little
C. It gets much worse: the query heads can no longer look for different things

*Answer it on the page to check your work.*

**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?

**Question.** 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)

*Answer it on the page to check your work.*

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](/learn/attention-backward/) 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](/math/#scientific) shows more examples.

**Question.** 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 it on the page to check your work.*

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 \times 4d^2$ |
| the FFN: $(d, d_{ff})$, then $(d_{ff}, d)$ | $(1, d)\,@\,(d, d_{ff})$, then $(1, d_{ff})\,@\,(d_{ff}, d)$ | $2 \times 2\,d\,d_{ff}$ |
| 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.

**Question.** 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 it on the page to check your work.*

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.

**Question.** 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?

*Answer it on the page to check your work.*

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

**Predict.** Before 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?

A. The arithmetic: 1.4 × 10¹⁰ FLOPs per token is a lot
B. Reading the weights from memory
C. Both take about the same time

*Answer it on the page to check your work.*

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:

$$
\text{tokens per second} \le \frac{\text{bandwidth}}{\text{weight bytes}}
$$

**Question.** 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 it on the page to check your work.*

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.

**Question.** 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 it on the page to check your work.*

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.

**Question.** 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 it on the page to check your work.*

*[Interactive lab: Serve — open the page to use it]*

**If you are stuck: 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.

**Question.** 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?

*Answer it on the page to check your work.*

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

**Question.** 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 it on the page to check your work.*

**Question.** 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 it on the page to check your work.*

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.

**Code question.** Write 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.

Fill in the blank (`____`):

```python
def serve(n_weights, kv_heads, d_k, layers, tokens, B,
          bandwidth, flops, bytes_per=2):  # bytes per number: weights and KV cache
    ____

# 7 billion weights, 8 K/V heads, 10¹² bytes/s, 10¹⁴ FLOPs/s
print(serve(7e9, 8, 128, 32, 0, 1, 1e12, 1e14))     # one user
print(serve(7e9, 8, 128, 32, 4000, 8, 1e12, 1e14))  # 8 users
```

*Answer it on the page to check your work.*

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.

## 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).
