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 (num_layers), (d_model), (ffw_dimension), (num_heads), (num_kv_heads), (qkv_dim), and (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:
Per-token KV cache size is:
Hence in int8:
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 = .
Note that:
where is total params in int8. Therefore the largest batch size we can fit in is:
For we get:
And if we let we get:
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 . Therefore to load into the MXU it takes:
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:
- What does ICI look like on a 4x4?
- What’s the roofline bound on tensor parallelism?
- 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:
Let , , and . Hence:
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:
During generation we will be memory bound instead of communication bound when using tensor parallelism when:
where in a TPU v5e. Hence:
Since for the maximum batch size is 7, then we get:
Therefore during generation we will be memory bound and a roofline bound will be:
We can shard the KV caches as:
where the mesh is defined as {X: 2, Y: 8}.
A rough per-step latency under this sharding is:
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 and with the above settings.
- How many total and activated parameters does it have? Activated means used by any given token.
- What batch size is needed to become FLOPs bound on TPU v5e?
- How large are its KV caches per token?
- How many FLOPs are involved in a forward pass with tokens?
Solution
Scan pages: 7-8
5.1
Let and , then:
5.2
From the approximation of the MLP part they use in this chapter, to be FLOPs bound we need :
5.3
KV cache size per token does not change.
5.4
Using the formula from Chapter 4:
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:
- What’s the HBM weight loading time for the above model on a TPU v5e 8x16 slice with , ? How much free HBM is available per TPU?
- What is the smallest slice we could fit our model on?
Solution
Scan pages: 8-10
6.1
Assuming we shard the MLPs as:
The attention params as:
and the shared input-output embedding as:
The total params per device is:
Hence the weight loading time in bf16:
As free space we have approximately:
6.2
Fit so that we can shard our expert copies, then we need to choose so that:
Here if we need a slice with axis sizes being multiple of 2, the smallest slice we could fit our model on is .
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:
Your goal is to work out and 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:
Note that given devices, can change based in how we allocate devices for the different axes , , . Therefore we need to find the setting where is minimum.
Solving , we get a candidate for a stationary point:
Since can only be positive in our case, we get:
Let to get:
Since at , we know that is a global minimum and becomes:
If we compare to traditional 3D model sharding where the condition to be compute bound is:
We see that with 2D weight stationary:
During inference however if we are communication bound in the MLP because of a small batch size, we see that: