Chapter 8: Serving LLaMA 3 on TPUs
Scaling Book Exercises – Chapter 8
Chapter: 8 – Serving LLaMA 3 on TPUs
Source: handwritten Onyx Boox A4 PDF
Note: Transcribed from handwritten solutions; book markdown used only for exercise statements and notation.
Remark: The following questions don’t appear numbered in the Scaling Book. I will number them in order of appearance as 0.X, where and .
Exercise 0.1 – LLaMA 3-70B KV cache size
Exercise statement
How large are LLaMA 3-70B’s KV caches per token? Assume we store them in int8. This determines how large our batch size can be on a given topology.
Solution
Scan pages: 1
Per-token KV cache size is:
where is sequence length. Hence, in int8:
Exercise 0.2 – Memory use and smallest slice
Exercise statement
Let’s say we want to serve L3 70B at batch size 32 and 8192 sequence length with everything (params and KVs) in int8. How much total memory will this use? What’s the smallest slice we could serve this on?
Solution
Scan pages: 1–2
Let be the number of parameters of the model, let and . Assume int8.
Assuming we serve this model in TPU v5e, the smallest possible slice is a since it gives us of aggregated HBM.
Exercise 0.3 – Decode latency and throughput
Exercise statement
At this batch size and quantization on a TPU v5e 4x2, roughly what latency would we expect per decode step? What throughput (tokens / sec / chip)? What about a 4x4? Assume we perform our FLOPs in bfloat16 and everything is fully sharded.
Solution
Scan pages: 2–3
At the MLP block is not compute bound. Therefore, using the general equation from chapter 6:
where is the total number of devices in the slice. Hence:
If instead we use a slice, we get:
Exercise 0.4 – Compute-bound batch sizes
Exercise statement
On TPU v5e, using bfloat16 weights and activations, how large do our batch sizes need to be for us to be compute-bound in our matmuls? What if we do int8 weights but perform our FLOPs in bfloat16? What about int8 weights with int8 FLOPs?
Solution
Scan pages: 3–4
For bfloat16 weights and activations:
If int8 weights but FLOPs in bfloat16, note that:
If int8 weights and int8 FLOPs:
Exercise 0.5 – Smallest serving topology by precision
Exercise statement
What is the smallest TPU v5e topology we could serve LLaMA 3-70B on using bfloat16, int8, and int4 (both KVs and parameters) with 8k context? You can think of KV caches as negligibly small for this one.
Solution
Scan pages: 4–5
Assume .
For bfloat16:
So we need at least a TPU v5e slice.
For int8:
So we need at least a TPU v5e slice.
For int4:
So we need at least a TPU v5e slice.
Exercise 0.6 – Generate-step latency at maximum batch size
Exercise statement
Assume we use the largest batch size that fits on these topologies. What latency could we expect for each generate step?
Solution
Scan pages: 5–6
For bfloat16 in a slice:
Note that cannot be , for otherwise our KV cache size will not fit in this topology. Hence we are memory bound, and choosing the max that satisfies the conditions just makes us closer to our upper bound. Therefore:
For int8, cannot be , hence as in the bfloat16 case:
For int4, the same argument as above holds, therefore:
Exercise 0.7 – Throughput per chip
Exercise statement
For each of these, what throughput per chip does this give us (in terms of queries / chip)? Assume our median decode length is 512 tokens.
Solution
Scan pages: 6–7
For each case we can use the condition we derived in Exercise 2 of chapter 7 to obtain the max that fits in each topology.
For int8 and a slice:
For bfloat16 and a slice:
For int4 and a slice:
Therefore:
Exercise 0.8 – Doubling the topology
Exercise statement
How would our peak throughput change if we doubled our topology for each of the above examples?
Solution
Scan pages: 8
If we double our topology, we can see that:
for all cases. Therefore, since:
we our peak throughput in each case.
Exercise 0.9 – Sharding on a TPU v5e 4x8
Exercise statement
Now let’s dig into the question of sharding. Let’s say we wanted to serve in bfloat16 on a TPU v5e 4x8. What sharding would we use for our model on a TPU v5e 4x8 during generation? Can we avoid being communication bound?
Solution
Scan pages: 8–9
The only sharding option we have is model parallelism. From chapter 7 we know that we will not be communication bound if:
where:
Therefore, for our setting:
So we can avoid being communication bound since in a TPU v5e slice.
The KV cache will be sharded along the head dimension and among the batch dimension as well:
where Mesh = {X: 4, Y: 8}.
Exercise 0.10 – Prefill time
Exercise statement
Assume we achieve a 40% FLOPs utilization during prefill. How long will a prefill of length 8192 take on 16 TPU v5e chips?
Solution
Scan pages: 9
We can use the formula:
Hence:
where we assumed bfloat16 FLOPs.
Exercise 0.11 – Decode completions and KV eviction
Exercise statement
Assume we have a median prefill length of 8192 tokens and a median decode length of 4096 tokens. Say we have a generate batch size of 32. On average how many sequences finish decoding per step? On average how many tokens are evicted from our KV cache each step?
Solution
Scan pages: 10
At batch size 32 we are memory bound during decoding, therefore:
Per prefill step I do:
So per prefill step I evict:
Exercise 0.12 – Prefill-to-generate server ratio
Exercise statement
Assume we do disaggregated serving with a median prefill length of 8192 and a median decode length of 512. Assume the prefill and generate latencies calculated above in bfloat16. What ratio of prefill:generate servers will you need to keep both fully saturated?
Solution
Scan pages: 11
Assuming each server produces:
Therefore we need to produce prefills at the same rate as the decoding servers:
So we need 7 more prefill servers than decoding servers.
Exercise 1 – LLaMA 3-405B forward-pass bounds
Exercise statement
How many FLOPs does each forward pass for LLaMA 3-405B use per-token? Assuming we’re FLOPs bound, what is a lower bound on a single forward pass on chips on TPU v5e? What if we’re comms bound? Ignore the fact that the model does not fit on a single chip.
Solution
Scan pages: 12
We can use the formula from this chapter:
Hence, per token:
If we are FLOPs bound, a lower bound on TPU v5e chips is:
If we are communication bound, a lower bound is:
Exercise 2 – LLaMA 3-8B serving memory
Exercise statement
Assume we want to serve LLaMA 3-8B with BS240 using int8 weights and int8 KV caches. How many bytes are used by (a) model parameters, (b) KV caches, and (c) peak working activations, roughly? What’s the smallest topology we can run this on?
Solution
Scan pages: 13
For LLaMA 3-8B we have , , , , , , and .
Assuming int8 for everything, we have:
For peak working activations we will assume:
Hence per token:
are used. The smallest topology is a single TPU v5e.
Exercise 3 – LLaMA 3-405B serving under a 15 ms limit
Exercise statement
How would you serve LLaMA 3-405B on TPU v5e? Assume int8 weights and bfloat16 FLOPs. Let’s say we have a firm limit of 15ms / token. What’s the highest-throughput configuration we could achieve? What is the theoretical minimum step time?
Solution
Scan pages: 14–17
For LLaMA 3-405B we have , , , , , , and . Assume int8 weights and bfloat16 FLOPs. Then we know that:
and:
Hence we need a TPU v5e slice topology with at least of aggregated HBM, assuming . Since HBM of TPU v5e is , we need:
Note that a lower bound in latency is:
Hence, if we need to have a latency of , we have to satisfy:
So we need more than 32 devices. We can consider an TPU v5e slice, and use model parallelism. During prefill we will be memory bound since:
However, we can improve our latency during generation.
Note that for a choice of we are memory bound if:
So then, in order to be memory bound, we need:
At this batch size we’re memory bound in the MLP. We can solve now for our max length that still fits:
However, we need to be smaller, for otherwise we will not hit the requirement.
Therefore, with and we have:
and a throughput of: