Chapter 6: Training LLaMA 3 on TPUs
Scaling Book Exercises – Chapter 6
Chapter: 6 – Training 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 FLOPs per token
Exercise statement
How many FLOPs does LLaMA-3 perform per token per training step ? This helps us determine how expensive the whole training process will be.
Solution
Scan pages: 1.
Let . We know that
Therefore per token we know its:
Exercise 0.2 – Total FLOPs for 15T tokens
Exercise statement
LLaMA 3 was trained for about 15 trillion tokens. How many FLOPs is that total ?
Solution
Scan pages: 1.
Let and , hence:
Exercise 0.3 – Training time on a full TPU v5p pod
Exercise statement
Let’s say we wanted to train on a full TPU v5p pod with 16x20x28 = 8960 chips. How long would this take to train at 40% MFU in bfloat16, assuming we are compute-bound?
Solution
Scan pages: 2.
Assume TPU v5p pod of chips. Let v5p bf16 FLOPs/s .
We know that FLOPs per chip is:
Let and . Therefore:
Exercise 0.4 – Minimum TPUs for the 4M-token batch
Exercise statement
LLaMA 3-70B was pretrained with a batch size of about 4M tokens. How many TPUs do we need at minimum to train with this batch size? You can assume bfloat16 parameters and float32 optimizer state, and that you checkpoint gradients 4 times per layer.
Solution
Scan pages: 2–3.
Assume we checkpoint the three big matmuls in the MLP part and , which has shape , so it occupies . Hence, combining our formula from exercise 5.2, we get that:
Therefore total memory required is:
where is number of model parameters. Let . Therefore for LLaMA 3-70B we get:
So assuming v5p TPU which has , we need at least:
TPUs.
Exercise 0.5 – Memory per chip on 8960 TPU v5p chips
Exercise statement
Under the same assumptions as the question above, if we use 8960 TPU v5p chips, how much memory will we use per-chip ?
Solution
Scan pages: 4.
Exercise 0.6 – FSDP alone without sequence/context parallelism
Exercise statement
Under the assumptions above, can we train our model with FSDP alone ? To start, let’s say we can’t do any sequence/context parallelism. This should be the first idea you have, since it’s simple and will introduce no extra communication if it works.
Solution
Scan pages: 4–6.
Modeling only the MLP part and using sharding strategy from [sequence parallelism, Shenggui Li et al.], both FSDP and sequence parallelism perform:
So without sequence parallelism we get:
Hence:
Therefore for being compute bound in this setting we need:
From the exercise we know that , therefore:
Since , we cannot train with FSDP alone.
Exercise 0.7 – FSDP with sequence/context parallelism
Exercise statement
Let’s relax the requirement of not doing any sequence sharding. If we allow ourselves to do FSDP over both the batch and sequence axes, can we train LLaMA 3-70B with only FSDP on 8960 chips?
Solution
Scan pages: 6–7.
Again modeling only the MLP part we can see that:
is equivalent to:
Since no attention involved we just have to reshape after finishing. Therefore with FSDP we will do:
Hence the condition to be compute bound is:
Since we cannot use FSDP if we want to be compute bound.
Exercise 0.8 – Mixed tensor parallelism and FSDP
Exercise statement
Now let’s look at mixed tensor parallelism and FSDP. Does there exist some combination that lets us remain compute-bound? What amount of FSDP and tensor parallelism should we do if so?
Solution
Scan pages: 7.
We require and to be , to be compute bound. Since and , if we choose , since is close to our , we will be compute bound.
Exercise 1 – Scaling LLaMA 70B to more chips
Exercise statement
Say we want to train LLaMA 3-70B on 4 pods with the same batch size. What parallelism scheme would we use? Would we be compute or communication bound? Roughly how long would it take to train? Make sure to use the correct roofline bound.
Solution
Scan pages: 8–11.
I need the per pod time to be compute bound, and we can use either FSDP or FSDP + TP. We can use data parallel over the DCN network and FSDP + TP. So let where and reshape to .
So the sharding strategy will be:
Note that , where . To be compute bound we need:
Assume and note that in the backward pass:
Therefore if , the first condition to be compute bound is satisfied if:
where DCN is the number of pods connected by DCN.
Now to satisfy condition 2 we need, as in the exercise before:
and choose integer close to as defined before. Since and , we get:
So if we choose and , we satisfy the conditions.
The time required for training is roughly, assuming MFU in bfloat16,
per 4M batch step. Therefore
Exercise 2.a – LLaMA 405B hyperparameters and FLOPs
Exercise statement
Using the LLaMA 3-405B config, write a table with all the key hyperparameters as above. How many total parameters does this model have? How many FLOPs per training step? How many FLOPs do we perform if we train for 15T tokens?
Solution
Scan pages: 11–12.
Using formulas from previous section since architecture does not change, just its hyperparams, we get:
Exercise 2.b – LLaMA 405B on 8 TPU v5p pods
Exercise statement
Assume we want to train on 8 TPU v5p pods. What parallelism scheme would we use? How long would training take? Would we be compute or comms bound?
Solution
Scan pages: 12–13.
Using our derivation from question 1, we can use that same sharding strategy and be compute bound if:
Choosing so that we are safely in the compute bound, we get , and choosing integer close to . In this case and .
A rough estimate for training time would be: