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

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 𝑀=70.4𝑒9. We know that

Transformer FLOPs ≈ 6 ⋅ 𝐵 𝑇 ⋅ 𝑀 .

Therefore per token we know its:

6 ⋅ 𝐵 𝑇 ⋅ 𝑀 𝐵 𝑇 = 6 ⋅ 𝑀 FLOPs = 422.4 𝑒 9 FLOPs .

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 𝑀=70.4𝑒9 and 𝐵𝑇=15𝑒12, hence:

Transformer FLOPs = 6 ⋅ 𝐵 𝑇 ⋅ 𝑀 = ( 6 ⋅ 70.4 ⋅ 15 ) ⋅ 10 21 = 6336 ⋅ 10 21 ≈ 6.3 𝑒 24 = 6.3 Yotta FLOPs .

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 16×20×28=8960 chips. Let v5p bf16 FLOPs/s =4.59𝑒14.

We know that FLOPs per chip is:

1 8960 ⋅ 6 ⋅ 𝐵 𝑇 ⋅ 𝑀 .

Let 𝑀=70.4𝑒9 and 𝐵𝑇=15𝑒12. Therefore:

Time = 6 ⋅ 𝐵 𝑇 ⋅ 𝑀 8960 ⋅ 0.4 ⋅ 4.59 𝑒 14 = 44.3 days .

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 𝐴=softmax(𝑄⋅𝐾𝑇)⋅𝑉, which has shape 𝐴[𝐵,𝑇,𝐾,𝐺,𝐻], so it occupies 2⋅𝐵𝑇𝐾𝐺𝐻=2𝐵𝑇𝐷. Hence, combining our formula from exercise 5.2, we get that:

Activation Checkpointing = 𝐿 ⋅ ( 4 𝐵 𝑇 𝐹 + 4 𝐵 𝑇 𝐷 ) .

Therefore total memory required is:

10 ⋅ 𝑀 + 𝐿 ⋅ ( 𝐵 𝑇 ) ⋅ ( 4 𝐹 + 4 𝐷 )

where 𝑀 is number of model parameters. Let 𝐵𝑇=4𝑒6. Therefore for LLaMA 3-70B we get:

10 ⋅ ( 70.4 𝑒 9 ) + 80 ⋅ ( 4 𝑒 6 ) ⋅ ( 4 ⋅ 28672 + 4 ⋅ 8192 ) ≈ 47.8 𝑒 12 = 47.8 TB .

So assuming v5p TPU which has HBM Capacity=96GB, we need at least:

47.8 𝑒 12 96 𝑒 9 ≈ 498

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.

Memory per chip = 47.8 𝑒 12 8960 bytes ≈ 5.3 𝑒 9 bytes = 5.3 GB .

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:

In [ 𝐵 𝑋 , 𝐿 𝑌 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 𝑋 , 𝐹 ] ⋅ 𝐹 𝑊 out [ 𝐹 , 𝐷 𝑋 ] .

So without sequence parallelism we get:

In [ 𝐵 𝑋 , 𝐿 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 𝑋 , 𝐹 ] ⋅ 𝐹 𝑊 out [ 𝐹 , 𝐷 𝑋 ] .

Hence:

𝑇 math = 2 ⋅ 2 ⋅ 𝐵 ⋅ 𝐿 ⋅ 𝐷 ⋅ 𝐹 𝑋 ⋅ 𝐶 𝑇 comms = 2 ⋅ 2 ⋅ 𝐷 ⋅ 𝐹 𝑊 ICI

Therefore for being compute bound in this setting we need:

4 𝐵 𝐿 𝐷 𝐹 4 𝑋 𝐷 𝐹 > 2550 ⇔ 𝐵 𝐿 𝑋 > 2550 ⇔ 𝐵 𝑋 > 2550 𝐿 .

From the exercise we know that 𝐿=4096, therefore:

𝐵 𝑋 > 0.62 .

Since 𝐵𝑋=0.11, 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:

In [ 𝐵 , 𝐿 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 , 𝐹 ] ⋅ 𝐹 𝑊 out [ 𝐹 , 𝐷 ]

is equivalent to:

In [ 𝐵 ⋅ 𝐿 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 , 𝐹 ] ⋅ 𝐹 𝑊 out [ 𝐹 , 𝐷 ] .

Since no attention involved we just have to reshape after finishing. Therefore with FSDP we will do:

In [ ( 𝐵 ⋅ 𝐿 ) 𝑋 , 𝐷 ] ⋅ 𝐷 𝑊 in [ 𝐷 𝑋 , 𝐹 ] ⋅ 𝑊 out [ 𝐹 , 𝐷 ] .

Hence the condition to be compute bound is:

𝐵 ⋅ 𝐿 𝑋 > 2550 .

Since 1024⋅40968960≈468 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 𝐵⋅𝐿𝑁>𝛼2𝑀𝑥𝑀𝑦⋅𝐹 and 𝑋 to be 𝑋opt=𝐵⋅𝐿𝐹⋅𝑀𝑥𝑀𝑦⋅𝑁, to be compute bound. Since 𝐵𝐿𝑁≈468 and 𝛼2𝑀𝑥𝑀𝑦⋅𝐹≈113, if we choose 𝑋=2240, since is close to our 𝑋opt, 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 𝐵=𝑁⋅𝐵local where 𝑁∈ℕ and reshape In[𝐵,𝐷] to In[𝑁,𝐵local,𝐷].

So the sharding strategy will be:

In [ 𝑁 DCN , 𝐵 local , 𝑋 , 𝐷 𝑌 ] ⋅ 𝐷 𝑊 in [ 𝐷 𝑋 , 𝐹 𝑌 ] ⋅ 𝐹 𝑊 out [ 𝐹 𝑌 , 𝐷 𝑋 ] .

Note that 𝑇=max(𝑇pod,𝑇comms-dcn), where 𝑇pod=max(𝑇math,𝑇comms). To be compute bound we need:

𝑇 comms-dcn < 𝑇 pod ∧ 𝑇 math > 𝑇 comms ⇒ 𝑇 math > 𝑇 comms-dcn . 𝑇 comms-dcn = 8 𝐷 𝐹 | 𝑋 | ⋅ | 𝑌 | .

Assume 𝑇math>𝑇comms and note that in the backward pass:

𝑇 math = 𝑁 DCN ⋅ 𝐵 local | 𝑋 | ⋅ 𝐷 ⋅ 𝐹 | 𝑌 | ⋅ 2 ⋅ 4 = 8 𝐵 𝐷 𝐹 DCN ⋅ | 𝑋 | ⋅ | 𝑌 | .

Therefore if 𝑇math>𝑇comms, the first condition to be compute bound is satisfied if:

8 𝐵 𝐷 𝐹 DCN ⋅ | 𝑋 | ⋅ | 𝑌 | ⋅ | 𝑋 | ⋅ | 𝑌 | 8 𝐷 𝐹 > 𝐶 𝑊 DCN = 73440 ⇒ 𝐵 DCN > 73440 ,

where DCN is the number of pods connected by DCN.

Now to satisfy condition 2 we need, as in the exercise before:

𝐵 local | 𝑋 | ⋅ | 𝑌 | > 𝛼 2 𝑀 𝑥 𝑀 𝑦 ⋅ 𝐹 ≈ 113 ,

and choose 𝑋 integer close to 𝑋opt as defined before. Since 𝐵=𝑁⋅𝐵local and 𝑁>DCN, we get:

𝐵 local > | 𝑋 | ⋅ | 𝑌 | ⋅ 113 ⇒ 𝑁 ⋅ 𝐵 local > 𝑁 ⋅ | 𝑋 | ⋅ | 𝑌 | ⋅ 113 ≫ DCN ⋅ | 𝑋 | ⋅ | 𝑌 | ⋅ 113 ⇒ 𝐵 > DCN ⋅ | 𝑋 | ⋅ | 𝑌 | ⋅ 113 .

So if we choose 𝐵=4𝑒6 and 𝑁=4, we satisfy the conditions.

The time required for training is roughly, assuming 40% MFU in bfloat16,

6 ⋅ 1 𝑒 6 ⋅ 70 𝑒 9 8960 ⋅ 0.4 ⋅ 4.59 𝑒 14 ≈ 255 ms

per 4M batch step. Therefore

15 𝑒 12 4 𝑒 6 ⋅ 255 ⋅ 10 − 3 ≈ 11 days .

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.

𝐿 = 126 , 𝐷 = 16384 , 𝐹 = 53248 , 𝑁 = 128 , 𝐾 = 8 , 𝐻 = 128 , 𝑉 = 128256 .

Using formulas from previous section since architecture does not change, just its hyperparams, we get:

Parameters = 405 𝑒 9 FLOPs per step = 6 ⋅ 405 𝑒 9 ≈ 2.4 TFLOPs FLOPs total = 6 ⋅ 15 𝑒 12 ⋅ 405 𝑒 9 = 3.645 𝑒 25 ≈ 36.45 Yotta FLOPs .

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:

𝐵 DCN > 73440 ∧ 𝐵 > DCN ⋅ | 𝑋 | ⋅ | 𝑌 | ⋅ 113 ∧ 𝑁 = DCN ∧ 𝐵 local = 𝐵 𝑁 .

Choosing 𝐵=16𝑒6 so that we are safely in the compute bound, we get 𝐵local=2𝑒6, and choosing 𝑋 integer close to 𝑋opt. In this case 𝑋=1120 and 𝑌=8.

A rough estimate for training time would be:

6 ⋅ 2 𝑒 6 ⋅ 405 𝑒 9 8960 ⋅ 0.4 ⋅ 4.59 𝑒 14 ⋅ 15 𝑒 12 16 𝑒 6 ≈ 32 days .