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 𝑋≥1.

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:

2 𝑆 𝐿 𝐾 𝐻 𝑆 = 2 𝐿 𝐾 𝐻

where 𝑆 is sequence length. Hence, in int8:

2 ⋅ 80 ⋅ 8 ⋅ 128 = 163840 Bytes ≈ 164 KB

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 𝐵=32 and 𝑆=8192. Assume int8.

Total Memory = 𝑀 + 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 = 70 𝑒 9 + 𝐵 𝑆 ⋅ 164 𝑒 3 ≈ 113 GB

Assuming we serve this model in TPU v5e, the smallest possible slice is a 4×2 since it gives us 128GB 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 𝐵=32 the MLP block is not compute bound. Therefore, using the general equation from chapter 6:

Step Time = 𝐵 × KV cache size BW ⋅ 𝑁 + 𝑀 BW ⋅ 𝑁

where 𝑁 is the total number of devices in the slice. Hence:

Step Time = 32 ⋅ 8192 ⋅ 164 𝑒 3 8.2 𝑒 11 ⋅ 8 + 70 𝑒 9 8.2 𝑒 11 ⋅ 8 ≈ 113 𝑒 9 8.2 𝑒 11 ⋅ 8 ≈ 0.017 s = 17 ms Tokens/s/chip = 𝐵 ⋅ 1 0.017 ⋅ 1 𝑁 = 𝐵 0.017 ⋅ 𝑁 = 32 0.017 ⋅ 8 = 235 tokens/s/chip

If instead we use a 4×4 slice, we get:

Step Time = 113 𝑒 9 8.2 𝑒 11 ⋅ 16 ≈ 0.008 s = 8 ms Tokens/s/chip = 32 0.008 ⋅ 16 = 250 tokens/s/chip

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:

Intensity(Matmul) = 𝐵 ≥ 1.97 𝑒 14 8.2 𝑒 11 ≈ 240

If int8 weights but FLOPs in bfloat16, note that:

Intensity(Matmul) = 2 𝐵 𝐷 𝐹 𝐵 𝐷 + 𝐷 𝐹 + 𝐵 𝐹 = 2 ⋅ 𝐵 𝐷 𝐹 𝐵 𝐷 + 𝐷 𝐹 + 𝐵 𝐹 ≈ 2 𝐵 ≥ 240 ⇔ 𝐵 ≥ 120

If int8 weights and int8 FLOPs:

Intensity(Matmul) = 2 𝐵 ≥ 2 ⋅ 1.97 𝑒 14 8.2 𝑒 11 ⇔ 𝐵 ≥ 240

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 𝐵=1.

For bfloat16:

Total Memory = 2 𝑀 + 2 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 = 2 ⋅ 70 𝑒 9 + 2 ⋅ 8192 ⋅ 164 𝑒 3 ≈ 143 GB

So we need at least a 4×4 TPU v5e slice.

For int8:

Total Memory = 𝑀 + 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 = 70 𝑒 9 + 8192 ⋅ 164 𝑒 3 ≈ 71 GB

So we need at least a 4×2 TPU v5e slice.

For int4:

Total Memory = 𝑀 2 + 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 2 ≈ 35.5 GB

So we need at least a 2×2 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 4×4 slice:

2 𝑀 + 2 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 4 × 4 ≤ 16 GB

Note that 𝐵 cannot be 𝐵≥240, 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:

bfloat16 Step Time = 16 𝑒 9 8.2 𝑒 11 ≈ 0.019 s = 19 ms

For int8, 𝐵 cannot be 𝐵≥240, hence as in the bfloat16 case:

int8 step time with max B = 19 ms

For int4, the same argument as above holds, therefore:

int4 step time with max B = 19 ms

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 2×4 slice:

𝐵 ≤ | 𝑋 | ⋅ HBM − 𝑀 2 𝑆 𝐿 𝐾 𝐻 = 8 ⋅ 16 − 70 𝑒 9 8192 ⋅ 164 𝑒 3 ≈ 42

For bfloat16 and a 4×4 slice:

𝐵 ≤ | 𝑋 | ⋅ HBM − 2 𝑀 2 𝑆 ⋅ 2 𝐿 𝐾 𝐻 ≈ 43

For int4 and a 2×2 slice:

𝐵 ≤ | 𝑋 | ⋅ HBM − 𝑀 2 𝑆 ⋅ 2 𝐿 𝐾 𝐻 2 ≈ 43

Therefore:

Queries/chip bfloat16, 4x4 slice = 1 0.014 ⋅ 512 ⋅ 42 ⋅ 1 16 ≈ 0.26 Queries/chip int8, 4x2 slice = 1 0.014 ⋅ 512 ⋅ 42 ⋅ 1 8 ≈ 0.53 Queries/chip int4, 2x2 slice = 1 0.014 ⋅ 512 ⋅ 42 ⋅ 1 4 ≈ 1.07

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:

𝐵 ≤ 2 ⋅ | 𝑋 | ⋅ HBM − 𝑀 𝑆 ⋅ 2 𝐿 𝐾 𝐻 ≈ 138

for all cases. Therefore, since:

138 42 ≈ 3.3

we 1.65× 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:

2 𝐷 𝐹 𝑁 𝑊 HBM > 2 𝐵 𝐷 2 𝑊 ICI ⇔ 2 𝐹 𝐵 𝛽 > 𝑁

where:

𝛽 = 𝑊 HBM 𝑊 ICI

Therefore, for our setting:

𝑁 < 2 𝐹 138 ⋅ 8 ≈ 52

So we can avoid being communication bound since 𝑁=32 in a 4×8 TPU v5e slice.

The KV cache will be sharded along the head dimension and among the batch dimension as well:

KV [ 𝐵 𝑋 , 2 , 𝑆 , 𝐿 , 𝐾 𝑌 , 𝐻 ]

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:

Total FLOPs = 2 𝑀 𝐵

Hence:

𝑡 = 2 𝑀 𝐵 16 ⋅ 1 𝐶 ⋅ 0.4 ≈ 0.9

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:

Step Time = 32 ⋅ 8192 ⋅ 164 𝑒 3 + 70 𝑒 9 𝑁 𝑊 HBM ≈ 0.008 s = 8 ms Sequence Time = Step Time × 4096 ≈ 35 s Decodings Per Second = 1 StepTime ⋅ 4096 ⋅ 32 ≈ 0.9 Prefills Per Second = 1 0.9 ≈ 1.1

Per prefill step I do:

0.9 ⋅ 32 35 = 0.82 decodings

So per prefill step I evict:

0.82 ⋅ 4096 ≈ 3358 tokens

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:

Prefills/s = 1 0.9 Decodings/s = 32 0.008 ⋅ 512 = 7.81

Therefore we need to produce prefills at the same rate as the decoding servers:

𝑁 ⋅ 1 0.9 = 7.81 ⇔ 𝑁 = 7.029

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:

Total FLOPs = 2 𝑀 𝐵

Hence, per token:

FLOPs/Token = 2 𝑀 = 2 ⋅ 405 𝑒 9 = 810 𝑒 9

If we are FLOPs bound, a lower bound on 𝑁 TPU v5e chips is:

𝑡 ≥ 2 𝑀 𝑁 𝐶 = 810 𝑒 9 𝑁 ⋅ 1.97 𝑒 14

If we are communication bound, a lower bound is:

𝑡 ≥ 2 𝐵 𝐷 𝑀 𝑋 ⋅ 𝑊 ICI = 2 ⋅ 𝐵 ⋅ 16384 𝑀 𝑋 ⋅ 9.0 𝑒 10

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 𝐿=32, 𝐷=4096, 𝐹=14336, 𝑁=32, 𝐾=8, 𝐻=128, and 𝑉=128256.

Assuming int8 for everything, we have:

𝑀 ≈ 8 𝑒 9 Bytes KV caches Size 𝑆 = 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 𝑆 = 240 ⋅ 2 ⋅ 32 ⋅ 8 ⋅ 128 ≈ 15.7 MB

For peak working activations we will assume:

In [ 𝐵 , 𝑆 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 , 𝐹 ] → Out [ 𝐵 , 𝑆 , 𝐹 ]

Hence per token:

𝐵 𝐹 = 240 ⋅ 14336 Bytes ≈ 3.4 MB

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 𝐿=126, 𝐷=16384, 𝐹=53248, 𝑁=128, 𝐾=8, 𝐻=128, and 𝑉=128256. Assume int8 weights and bfloat16 FLOPs. Then we know that:

KV cache size/token = 2 𝐿 𝐾 𝐻 = 2 ⋅ 126 ⋅ 8 ⋅ 128 ≈ 258 𝑒 3 Bytes = 258 KB

and:

Total Memory = 𝑀 + 𝐵 𝑆 ⋅ 2 𝐿 𝐾 𝐻 = 405 𝑒 9 + 𝐵 𝑆 ⋅ 258 𝑒 3

Hence we need a TPU v5e slice topology with at least 408GB of aggregated HBM, assuming 𝐵=1. Since HBM of TPU v5e is 16GB, we need:

| 𝑋 | ⋅ 16 GB ≥ 408 GB ⇔ | 𝑋 | ≥ 25.5

Note that a lower bound in latency is:

405 𝑒 9 | 𝑋 | ⋅ 𝑊 HBM ≤ Step Time

Hence, if we need to have a latency of ≤15ms, we have to satisfy:

405 𝑒 9 0.015 ⋅ 8.2 𝑒 11 ≤ | 𝑋 | ⇔ | 𝑋 | ≥ 32.9

So we need more than 32 devices. We can consider an 8×8 TPU v5e slice, and use model parallelism. During prefill we will be memory bound since:

𝐹 2188 = 24 ≱ | 𝑋 | = 64

However, we can improve our latency during generation.

Note that for a choice of 𝐵 we are memory bound if:

𝑇 HBM comms ( int8 ) = 𝐷 𝐹 | 𝑋 | ⋅ 𝑊 HBM 𝑇 ICI comms ( bfloat16 ) = 2 𝐵 𝐷 𝑊 ICI 𝐷 𝐹 | 𝑋 | ⋅ 𝑊 HBM > 2 𝐵 𝐷 𝑊 ICI ⇔ 𝐹 ⋅ 𝑊 ICI 2 ⋅ | 𝑋 | ⋅ 𝑊 HBM > 𝐵

So then, in order to be memory bound, we need:

𝐵 ≤ 52.0

At this batch size we’re memory bound in the MLP. We can solve now for our max 𝑆 length that still fits:

𝑆 ≤ | 𝑋 | ⋅ HBM − 𝑀 𝐵 ⋅ 2 𝐿 𝐾 𝐻 ≈ 46130

However, we need 𝑆 to be smaller, for otherwise we will not hit the 15ms requirement.

Step time = 52 ⋅ 𝑆 ⋅ 258 𝑒 3 + 405 𝑒 9 64 ⋅ 8.2 𝑒 11 ≤ 0.015 𝑆 ≤ 64 ⋅ 8.2 𝑒 11 ⋅ 0.015 − 405 𝑒 9 52 ⋅ 258 𝑒 3 ≈ 28488

Therefore, with 𝐵=52 and 𝑆=28488 we have:

Step Time = 15 ms

and a throughput of:

Tokens/s = 𝐵 0.015 = 52 0.015 = 3466