Chapter 7: All About Transformer Inference

Scaling Book Exercises – Chapter 7

Chapter: 7 – All About Transformer Inference
Source: handwritten Onyx Boox A4 PDF
Note: Transcribed from handwritten solutions; book markdown used only for exercise statements and notation.

Shared model setup for Exercises 1-7

The chapter uses an invented model based on LLaMA-2 13B with 𝐿=64 (num_layers), 𝐷=4096 (d_model), 𝐹=16384 (ffw_dimension), 𝑁=32 (num_heads), 𝐾=8 (num_kv_heads), 𝐻=256 (qkv_dim), and 𝑉=32128 (num_embeddings).

Exercise 1 – Parameter count and int8 KV cache

Exercise statement

How many parameters does the above model have? How large are its KV caches per token in int8? You can assume we share the input and output projection matrices.

Solution

Scan pages: 1

From previous chapters we know that number of params is:

Params = ( 3 𝐷 𝐹 + 2 𝐷 ( 𝑁 + 𝐾 ) ⋅ 𝐻 ) ⋅ 𝐿 + 𝐷 𝑉 = ( 3 ⋅ 4096 ⋅ ( 4 ⋅ 4096 ) + 2 ⋅ 4096 ⋅ 40 ⋅ 256 ) ⋅ 64 + 4096 ⋅ 32128 = 18 385 207 246 ≈ 18 ⋅ 10 9 = 18 Billion .

Per-token KV cache size is:

2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 𝑆 = 2 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 .

Hence in int8:

2 ⋅ 64 ⋅ 8 ⋅ 256 = 262144 bytes ≈ 262 KB .

Exercise 2 – Largest batch size at 128K context

Exercise statement

Say we want to serve this model on a TPUv5e 4x4 slice and can fully shard our KV cache over this topology. What’s the largest batch size we can fit, assuming we use int8 for everything and want to support 128k sequences? What if we dropped the number of KV heads to 1?

Solution

Scan pages: 2-3

Assume TPU v5e 4x4 slice where v5e HBM = 16GB.

Note that:

Space per shard = 𝑀 + 𝐵 2 𝑆 𝐿 𝐾 𝐻 | 𝑋 | . 𝑀 | 𝑋 | + 𝐵 ⋅ 2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 | 𝑋 | ≤ HBM .

where 𝑀 is total params in int8. Therefore the largest batch size we can fit in is:

𝑀 | 𝑋 | + 𝐵 ⋅ 2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 | 𝑋 | ≤ HBM ⇔ 𝑀 | 𝑋 | − HBM ≤ − 𝐵 ⋅ 2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 | 𝑋 | ⇔ | 𝑋 | ⋅ HBM − 𝑀 | 𝑋 | ≥ 𝐵 ⋅ 2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 | 𝑋 | ⇔ | 𝑋 | ⋅ HBM − 𝑀 2 ⋅ 𝑆 ⋅ 𝐿 ⋅ 𝐾 ⋅ 𝐻 ≥ 𝐵 .

For 𝑆=128K we get:

𝐵 ≤ 16 ⋅ ( 16 ⋅ 10 9 ) − 18 ⋅ 10 9 2 ⋅ 128 ⋅ 10 3 ⋅ 64 ⋅ 8 ⋅ 256 = 7.04 ≈ 7 .

And if we let 𝐾=1 we get:

𝐵 ≤ 16 ⋅ 16 ⋅ 10 9 − 18 ⋅ 10 9 2 ⋅ 128 ⋅ 10 3 ⋅ 64 ⋅ 1 ⋅ 256 ≈ 56 .

Exercise 3 – Parameter-loading latency

Exercise statement

How long does it take to load all the parameters into the MXU from HBM assuming they’re fully sharded on a TPU v5e 4x4 slice? Assume int8 parameters. This is a good lower bound on the per-step latency.

Solution

Scan pages: 3

Assume TPU v5e 4x4 slice where v5e BW=8.2⋅1011. Therefore to load 𝑀|𝑋| into the MXU it takes:

𝑀 | 𝑋 | ⋅ 1 BW = 18 ⋅ 10 9 16 ⋅ 1 8.2 ⋅ 10 11 ≈ 0.0013 s = 1.3 ms .

Exercise 4 – Prefill and generation sharding

Exercise statement

Let’s say we want to serve this model on a TPUv5e 4x4 slice using int8 FLOPs and parameters/activations. How would we shard it for both prefill and decode? Hint: maybe answer these questions first:

  1. What does ICI look like on a 4x4?
  2. What’s the roofline bound on tensor parallelism?
  3. How can we shard the KV caches?

For this sharding, what is the rough per-step latency for generation?

Solution

Scan pages: 4-6

Assume TPU v5e 4x4 slice. With this slice we don’t have wraparound links. To shard it both for prefill and generation, since during generation FSDP does not help, we could use tensor parallelism sharded over all the devices for both prefill and generation.

For prefill we are compute bound when:

𝑌 < 𝐹 ⋅ 𝑊 ICI 𝐶 ⋅ 2 .

Let 𝐹=4⋅4096, 𝑊ICI=9⋅1010, and 𝐶=3.94⋅1014. Hence:

𝑌 < 16384 ⋅ 9 ⋅ 10 10 3.94 ⋅ 10 14 ⋅ 2 ≈ 7.5 .

So we cannot do tensor parallelism with 16 devices and be compute bound.

We could instead do TP and CP with 4 way TP and 4 way CP, to be compute bound. So then a roofline bound will be:

𝑡 = 4 ⋅ 𝐵 𝐷 𝐹 𝑋 𝑌 𝐶 , where 𝐵 = 7 ⋅ 128 K .

During generation we will be memory bound instead of communication bound when using tensor parallelism when:

𝑌 < 𝐹 ⋅ 1 𝐵 ⋅ 𝑊 ICI 𝑊 HBM .

where 𝑊ICI𝑊HBM≈18 in a TPU v5e. Hence:

𝑌 < 4 ⋅ 4096 ⋅ 1 𝐵 ⋅ 1 8 = 2048 𝐵 .

Since for 𝑆=128K the maximum batch size is 7, then we get:

𝑌 < 2048 7 = 292 .

Therefore during generation we will be memory bound and a roofline bound will be:

𝑡 = 2 𝐷 𝐹 𝑋 ⋅ 𝑊 HBM .

We can shard the KV caches as:

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

where the mesh is defined as {X: 2, Y: 8}.

A rough per-step latency under this sharding is:

𝐵 × KV cache size + Parameter size 16 ⋅ 𝑊 HBM = 7 × ( 2 ⋅ 128 ⋅ 10 3 ⋅ 64 ⋅ 8 ⋅ 256 ) + 18 ⋅ 10 9 16 ⋅ 8.2 ⋅ 10 11 ≈ 0.02 s .

Exercise 5 – MoE parameter and FLOP accounting

Exercise statement

Let’s pretend the above model is actually an MoE. An MoE model is effectively a dense model with 𝐸 copies of the FFW block. Each token passes through 𝑘 of the FFW blocks and these 𝑘 are averaged to produce the output. Let’s use 𝐸=16 and 𝑘=2 with the above settings.

  1. How many total and activated parameters does it have? Activated means used by any given token.
  2. What batch size is needed to become FLOPs bound on TPU v5e?
  3. How large are its KV caches per token?
  4. How many FLOPs are involved in a forward pass with 𝑇 tokens?

Solution

Scan pages: 7-8

5.1

Let 𝐸=16 and 𝑘=2, then:

Total Params = ( 3 𝐸 𝐷 𝐹 + 2 𝐷 ( 𝑁 + 𝐾 ) ⋅ 𝐻 ) ⋅ 𝐿 + 𝐷 𝑉 ≈ 212 Billion . Activated Params = ( 3 𝑘 𝐷 𝐹 + 2 𝐷 ( 𝑁 + 𝐾 ) ⋅ 𝐻 ) ⋅ 𝐿 + 𝐷 𝑉 ≈ 31 Billion .

5.2

From the approximation of the MLP part they use in this chapter, to be FLOPs bound we need 𝑇math>𝑇comms:

2 ⋅ 𝐵 ⋅ 3 𝑘 𝐷 𝐹 𝐶 > 2 ⋅ 3 ⋅ 𝐸 𝐷 𝐹 BW ⇔ 2 ⋅ 𝐵 ⋅ 3 𝑘 𝐷 𝐹 2 ⋅ 3 ⋅ 𝐸 𝐷 𝐹 > 𝐶 BW ⇔ 𝐵 ⋅ 𝑘 𝐸 > 𝐶 BW ⇔ 𝐵 > 𝐸 𝑘 ⋅ 𝐶 BW .

5.3

KV cache size per token does not change.

5.4

Using the formula from Chapter 4:

FLOPs = 6 ⋅ 𝐵 𝑇 ( 3 𝑘 𝐷 𝐹 + 2 ⋅ 𝐷 ( 𝑁 + 𝐾 ) ⋅ 𝐻 ) ⋅ 𝐿 .

Exercise 6 – Expert sharding and minimum slice

Exercise statement

With MoEs, we can do “expert sharding”, where we split our experts across one axis of our mesh. In our standard notation, our first FFW weight has shape [E, D, F] and we shard it as [𝐸𝑍,𝐷𝑋,𝐹𝑌] where 𝑋 is only used during training as our FSDP dimension. Let’s say we want to do inference on a TPU v5e:

  1. What’s the HBM weight loading time for the above model on a TPU v5e 8x16 slice with 𝑌=8, 𝑍=16? How much free HBM is available per TPU?
  2. What is the smallest slice we could fit our model on?

Solution

Scan pages: 8-10

6.1

Assuming we shard the MLPs as:

𝑊 in1 [ 𝐸 𝑍 , 𝐷 , 𝐹 𝑌 ] , 𝑊 in2 [ 𝐸 𝑍 , 𝐷 , 𝐹 𝑌 ] , 𝑊 out [ 𝐸 𝑍 , 𝐹 𝑌 , 𝐷 ] .

The attention params as:

𝑊 𝑞 [ 𝐷 , 𝑁 𝑌 , 𝐻 ] , 𝑊 𝑘 [ 𝐷 , 𝐾 𝑌 , 𝐻 ] , 𝑊 𝑣 [ 𝐷 , 𝐾 𝑌 , 𝐻 ] , 𝑊 𝑜 [ 𝑁 𝑌 , 𝐻 , 𝐷 ] .

and the shared input-output embedding as:

𝑉 [ 𝑉 𝑌 , 𝐷 ] .

The total params per device is:

Params Per Device = ( 3 𝐸 | 𝑍 | 𝐷 𝐹 | 𝑌 | + 2 𝐷 | 𝑌 | ( 𝑁 + 𝐾 ) 𝐻 ) ⋅ 𝐿 + 𝐷 𝑉 | 𝑌 | ≈ 2.3 Billion .

Hence the weight loading time in bf16:

2 ⋅ 2.3 ⋅ 10 9 8.2 ⋅ 10 11 ≈ 0.005 s = 5 ms .

As free space we have approximately:

16 GB − 4.6 GB = 11.4 GB .

6.2

Fit 𝑍=16 so that we can shard our expert copies, then we need to choose 𝑌 so that:

2 ⋅ 1 | 𝑌 | ( ( 3 𝐷 𝐹 + 2 𝐷 ( 𝑁 + 𝐾 ) ⋅ 𝐻 ) ⋅ 𝐿 + 𝐷 𝑉 ) ≤ 16 ⇔ 2 ⋅ 18 ⋅ 10 9 16 ≤ | 𝑌 | ⇔ | 𝑌 | ≥ 2.25 .

Here if we need a slice with axis sizes being multiple of 2, the smallest slice we could fit our model on is 4×16.

Exercise 7 – 2D model sharding

Exercise statement

Here we’ll work through the math of what the ESTI paper calls 2D weight-stationary sharding. We describe this briefly in Appendix B, but try doing this problem first to see if you can work out the math. The basic idea of 2D weight stationary sharding is to shard our weights along both the 𝐷 and 𝐹 axes so that each chunk is roughly square. This reduces the comms load and allows us to scale slightly farther.

Here’s the algorithm for 2D weight stationary:

  1. In[𝐵,𝐷𝑋]=AllGather𝑌𝑍(In[𝐵,𝐷𝑋𝑌𝑍])
  2. Tmp[𝐵,𝐹𝑌𝑍]{𝑈𝑋}=In[𝐵,𝐷𝑋]⋅𝐷𝑊in[𝐷𝑋,𝐹𝑌𝑍]
  3. Tmp[𝐵,𝐹𝑌𝑍]=AllReduce𝑋(Tmp[𝐵,𝐹𝑌𝑍]{𝑈𝑋})
  4. Out[𝐵,𝐷𝑋]{𝑈𝑌𝑍}=Tmp[𝐵,𝐹𝑌𝑍]⋅𝐹𝑊out[𝐹𝑌𝑍,𝐷𝑋]
  5. Out[𝐵,𝐷𝑋𝑌𝑍]=ReduceScatter𝑌𝑍(Out[𝐵,𝐷𝑋]{𝑈𝑌𝑍})

Your goal is to work out 𝑇math and 𝑇comms for this algorithm and find when it will outperform traditional 3D model sharding?

Solution

Scan pages: 10-14

Let 𝑁=|𝑋|⋅|𝑌|⋅|𝑍| be the total number of devices in the slice. From the algorithm we can see that:

𝑇 comms = 2 𝐵 𝐷 | 𝑋 | ⋅ 1 2 ⋅ 𝑊 ICI + 2 ⋅ 2 ⋅ 𝐵 𝐹 | 𝑌 | ⋅ | 𝑍 | ⋅ 1 𝑊 ICI + 2 𝐵 𝐷 | 𝑋 | ⋅ 1 2 ⋅ 𝑊 ICI = 2 𝐵 𝐷 | 𝑋 | ⋅ 𝑊 ICI + 4 𝐵 𝐹 ⋅ | 𝑋 | 𝑁 ⋅ 𝑊 ICI = 1 𝑊 ICI ( 2 𝐵 𝐷 | 𝑋 | + 4 𝐵 𝐹 ⋅ | 𝑋 | 𝑁 ) . 𝑇 math = 2 ⋅ 𝐵 ⋅ 𝐷 | 𝑋 | ⋅ 𝐹 | 𝑌 | ⋅ | 𝑍 | ⋅ 1 𝐶 + 2 ⋅ 𝐵 ⋅ 𝐹 | 𝑌 | ⋅ | 𝑍 | ⋅ 𝐷 | 𝑋 | ⋅ 1 𝐶 = 4 𝐵 𝐷 𝐹 | 𝑋 | ⋅ | 𝑌 | ⋅ | 𝑍 | ⋅ 1 𝐶 = 4 𝐵 𝐷 𝐹 𝑁 ⋅ 𝐶 .

Note that given 𝑁 devices, 𝑇comms can change based in how we allocate devices for the different axes 𝑋, 𝑌, 𝑍. Therefore we need to find the setting where 𝑇comms is minimum.

𝑇 comms ( | 𝑋 | ) = 1 𝑊 ICI ( 2 𝐵 𝐷 | 𝑋 | + 4 𝐵 𝐹 ⋅ | 𝑋 | 𝑁 ) . 𝜕 𝑇 comms 𝜕 | 𝑋 | = 1 𝑊 ICI ( 2 𝐵 𝐷 ⋅ ( − 1 ) ⋅ 1 | 𝑋 | 2 + 4 𝐵 𝐹 𝑁 ) .

Solving 𝜕𝑇comms𝜕|𝑋|=0, we get a candidate for a stationary point:

4 𝐵 𝐹 𝑁 = 2 𝐵 𝐷 | 𝑋 | 2 ⇔ 𝑁 4 𝐵 𝐹 ⋅ 2 𝐵 𝐷 = | 𝑋 | 2 ⇔ | 𝑋 | 2 = 𝐷 2 𝐹 ⋅ 𝑁 .

Since |𝑋| can only be positive in our case, we get:

𝑋 ∗ = 𝐷 2 𝐹 ⋅ 𝑁 .

Let 𝐹=4𝐷 to get:

𝑋 ∗ = 𝑁 8 .

Since 𝜕2𝑇comms𝜕|𝑋|2>0 at |𝑋|=𝑋∗, we know that 𝑋∗ is a global minimum and 𝑇comms becomes:

𝑇 comms = 1 𝑊 ICI ( 2 𝐵 𝐷 8 𝑁 + 4 𝐵 𝐹 𝑁 𝑁 8 ) = 1 𝑊 ICI ( 2 𝐵 𝐷 8 𝑁 + 4 𝐵 𝐹 𝑁 8 ) = 1 𝑊 ICI ( 16 𝐵 𝐷 8 𝑁 + 4 𝐵 𝐹 𝑁 8 ) = 1 8 ⋅ 𝑊 ICI ( 4 𝐵 ⋅ ( 4 𝐷 + 𝐹 ) 𝑁 ) = 1 8 ⋅ 𝑊 ICI ( 4 𝐵 ⋅ 8 𝐷 𝑁 ) = 1 8 ⋅ 𝑊 ICI ⋅ 32 𝐵 𝐷 𝑁 . 𝑇 math 𝑇 comms > 1 ⇔ 4 𝐵 𝐷 𝐹 𝑁 𝐶 ⋅ 8 𝑊 ICI 𝑁 32 𝐵 𝐷 > 1 ⇔ 4 8 𝑊 ICI 𝑁 𝐹 32 𝑁 𝐶 > 1 ⇔ 4 8 32 ⋅ 1 𝑁 ⋅ 𝑊 ICI 𝐹 𝐶 > 1 ⇔ 16 ⋅ 8 32 2 ⋅ 1 𝑁 ⋅ 𝑊 ICI 2 𝐹 2 𝐶 2 > 1 ⇔ 16 ⋅ 8 32 2 ⋅ 𝐹 2 𝛼 2 > 𝑁 ⇔ 1 8 ( 𝐹 𝛼 ) 2 > 𝑁 .

If we compare to traditional 3D model sharding where the condition to be compute bound is:

𝑀 𝑦 ⋅ 𝐹 𝛼 > 𝑁 .

We see that with 2D weight stationary:

1 8 ( 𝐹 𝛼 ) 2 𝑀 𝑥 𝑦 𝑧 ( 𝐹 𝛼 ) = 1 8 𝑀 𝑥 𝑦 𝑧 ⋅ 𝐹 𝛼 .

During inference however if we are communication bound in the MLP because of a small batch size, we see that:

𝑇 comms-2D 𝑇 comms-3D = 32 8 ⋅ 𝐵 𝐷 𝑊 ICI 𝑁 ⋅ 𝑊 ICI 𝑀 𝑥 𝑦 𝑧 4 𝐵 𝐷 = 32 8 ⋅ 4 ⋅ 𝑀 𝑥 𝑦 𝑧 𝑁 = 2 ⋅ 2 ⋅ 𝑀 𝑥 𝑦 𝑧 𝑁 .