Chapter 5: How to Parallelize a Transformer for Training

Scaling Book Exercises – Chapter 5

Chapter: 5 – How to Parallelize a Transformer for Training
Source: handwritten Onyx Boox A4 PDF
Note: Transcribed from handwritten solutions; book markdown used only for exercise statements and notation.

Shared setup from the book

Let’s use LLaMA-2 13B as a basic model for this section. Here are the model details:

hyperparam value
𝐿 40
𝐷 5,120
𝐹 13824
𝑁 40
𝐾 40
𝐻 128
𝑉 32,000

LLaMA-2 has separate embedding and output matrices and a gated MLP block.

Exercise 1 – LLaMA-2 13B parameter count

Exercise statement

How many parameters does LLaMA-2 13B have (I know that’s silly but do the math)? Note that, as in Transformer Math, LLaMA-3 has 3 big FFW matrices, two up-projection and one down-projection. We ignored the two “gating” einsum matrices in this section, but they behave the same as 𝑊in in this section.

Solution

Scan pages: 1

From LLAMA2 we know that 𝑁=𝐾. Therefore:

MLP Params = 3 𝐷 𝐹 per layer Attention Params = 4 𝐷 𝑁 𝐻 per layer Embeddings = 2 𝐷 𝑉 overall .

In total we have (3𝐷𝐹+4𝐷𝑁𝐻)⋅𝐿+2𝐷𝑉 parameters. Let 𝐿=40, 𝑁𝐻=𝐷, 𝐷=5120, 𝐹=13824, and 𝑉=32,000. Hence:

( 3 ⋅ 5120 ⋅ 13824 + 4 ( 5120 ) 2 ) ⋅ 40 + 2 ⋅ 5120 ⋅ 32 𝑒 3 = 13 , 015 , 449 , 600 ≈ 13 ⋅ 10 9 = 13 Billion .

Exercise 2 – Memory for BS=16M training

Exercise statement

Let’s assume we’re training with BS=16M tokens and using Adam. Ignoring parallelism for a moment, how much total memory is used by the model’s parameters, optimizer state, and activations? Assume we store the parameters in bf16 and the optimizer state in fp32 and checkpoint activations three times per layer (after the three big matmuls).

Solution

Scan pages: 1-2

Model Params =2⋅13𝑒9=26𝑒9 Bytes=26GB.

Optimizer State:

4 ⋅ 13 𝑒 9 + 4 ⋅ 13 𝑒 9 = 8 ⋅ 13 𝑒 9 = 104 GB .

Activation Checkpointing:

𝐿 ⋅ ( 2 𝐵 𝑇 𝐹 + 2 𝐵 𝑇 𝐹 + 2 𝐵 𝑇 𝐷 ) = 40 ⋅ ( 4 𝐵 𝑇 𝐹 + 2 𝐵 𝑇 𝐷 ) = 40 ⋅ ( 𝐵 𝑇 ) ⋅ ( 4 𝐹 + 2 𝐷 ) = 40 ⋅ ( 16 𝑒 6 ) ⋅ ( 4 ⋅ 13824 + 2 ⋅ 5120 ) ≈ 42 TB .

Exercise 3 – 32k training on TPU v5p 16x16x16

Exercise statement

Assume we want to train with 32k sequence length and a total batch size of 3M tokens on a TPUv5p 16x16x16 slice. Assume we want to use bfloat16 weights and a float32 optimizer, as above.

  1. Can we use pure data parallelism? Why or why not?
  2. Can we use pure FSDP? Why or why not? With pure FSDP, how much memory will be used per device (assume we do gradient checkpointing only after the 3 big FFW matrices).
  3. Can we use mixed FSDP + tensor parallelism? Why or why not? If so, what should 𝑋 and 𝑌 be? How much memory will be stored per device? Using only roofline FLOPs estimates and ignoring attention, how long will each training step take at 40% MFU?

Solution

Scan pages: 2-11

Let:

HBM v5p = 96 GB bf16 FLOPs/s v5p ≈ 4.59 𝑒 14 𝑊 ICI v5p ≈ 1.8 𝑒 11 .
3.1

We cannot use data parallelism since a lower bound in memory per chip we need is 130GB. (In fact we need more if we consider the activation checkpointing.)

3.2

Assume FSDP shards the whole architecture as follows. For the attention params:

𝑊 𝑞 [ 𝐷 𝑋 , 𝑁 , 𝐻 ] 𝑊 𝑘 [ 𝐷 𝑋 , 𝑁 , 𝐻 ] 𝑊 𝑣 [ 𝐷 𝑋 , 𝑁 , 𝐻 ] 𝑊 𝑜 [ 𝑁 , 𝐻 , 𝐷 𝑋 ]

and for the MLP params as usual:

𝑊 in1 [ 𝐷 𝑋 , 𝐹 ] 𝑊 in2 [ 𝐷 𝑋 , 𝐹 ] 𝑊 out [ 𝐹 , 𝐷 𝑋 ]

and for input output embeddings as:

𝑊 embed [ 𝑉 , 𝐷 𝑋 ] 𝑊 outembed [ 𝐷 𝑋 , 𝑉 ]

Therefore the total params per device is:

( 3 𝐷 | 𝑋 | 𝐹 + 4 𝐷 | 𝑋 | 𝐻 𝑁 ) ⋅ 𝐿 + 2 𝐷 | 𝑋 | 𝑉 = 1 | 𝑋 | ⋅ ( ( 3 𝐷 𝐹 + 4 𝐷 𝐻 𝑁 ) ⋅ 𝐿 + 2 𝐷 𝑉 ) ≈ Model Params | 𝑋 | .

For LLAMA2 13B therefore:

Device Params = 13 ⋅ 10 9 16 3 ≈ 3.2 ⋅ 10 6 = 3.2 Million .

Assuming optimizer states also sharded, we get that each device needs at least:

10 ⋅ ( 3.2 ⋅ 10 6 ) = 32 ⋅ 10 6 = 32 MB .

If we also take into account the activation checkpointings which are of shapes:

[ 𝐵 𝑋 , 𝑇 , 𝐹 ] [ 𝐵 𝑋 , 𝑇 , 𝐹 ] [ 𝐵 𝑋 , 𝑇 , 𝐷 ]

we get:

Checkpoints per device = 𝐿 ⋅ ( 2 𝐵 | 𝑋 | ⋅ 𝑇 𝐷 + 4 𝐵 | 𝑋 | 𝑇 𝐹 ) Bytes = 1 | 𝑋 | ⋅ ( 𝐿 ⋅ ( 𝐵 𝑇 ) ⋅ ( 4 𝐹 + 2 𝐷 ) ) Bytes .

Let 𝐵𝑇=3⋅106, then:

1 16 3 ⋅ ( 40 ⋅ ( 3 ⋅ 10 6 ) ⋅ ( 4 ⋅ 13824 + 2 ⋅ 5120 ) ) ≈ 1.9 GB .

So memory used per device is 32MB+1.96GB.

We could use FSDP since with FSDP sharding everything fits in memory, but we would be communication bound since we are compute bound iff:

𝐵 | 𝑋 | > 2550 𝑀 𝑋 ⇔ 𝑀 𝑋 ⋅ 𝐵 2550 > | 𝑋 | .

Let 𝐵=3⋅106 and 𝑀𝑋=3, therefore:

9 ⋅ 10 6 2550 ≈ 3530 > | 𝑋 | ,

and in our slice |𝑋|=4046.

3.3

Assume we shard the attention matrices as:

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

and for the MLP params as:

𝑊 in1 [ 𝐷 𝑋 , 𝐹 𝑌 ] 𝑊 in2 [ 𝐷 𝑋 , 𝐹 𝑌 ] 𝑊 out [ 𝐹 𝑌 , 𝐷 𝑋 ]

and for input and output embeddings as:

𝑉 embed [ 𝑉 𝑌 , 𝐷 𝑋 ] 𝑉 outembed [ 𝐷 𝑋 , 𝑉 𝑌 ]

Therefore total params per device is:

( 3 𝐷 | 𝑋 | | 𝑌 | 𝐹 + 4 𝐷 | 𝑋 | ⋅ 𝐻 𝑁 | 𝑌 | ) ⋅ 𝐿 + 2 𝐷 | 𝑋 | | 𝑌 | 𝑉 = 1 | 𝑋 | | 𝑌 | ⋅ ( 𝐿 ⋅ ( 3 𝐷 𝐹 + 4 𝐷 𝐻 𝑁 ) + 2 𝐷 𝑉 ) = Model Params | 𝑋 | | 𝑌 | .

For LLAMA2 13B we get:

Params per device = 13 ⋅ 10 9 ( 16 ) 3 ≈ 3.2 Million .

If we take into account activation state then we get:

10 ⋅ ( 3.2 ⋅ 10 6 ) = 32 MB .

Activation checkpoints will be sharded with shapes:

[ 𝐵 𝑋 , 𝑇 , 𝐹 𝑌 ] [ 𝐵 𝑋 , 𝑇 , 𝐹 𝑌 ] [ 𝐵 𝑋 , 𝑇 , 𝐷 𝑌 ]

so we get:

Checkpoints per device = 𝐿 ⋅ ( 2 𝐵 | 𝑋 | ⋅ 𝑇 𝐷 | 𝑌 | + 4 𝐵 | 𝑋 | 𝑇 𝐹 | 𝑌 | ) Bytes = 1 | 𝑋 | | 𝑌 | ⋅ ( 𝐿 ⋅ ( 𝐵 𝑇 ) ⋅ ( 4 𝐹 + 2 𝐷 ) ) Bytes .

Let 𝐵𝑇=3⋅106, then:

1 16 3 ⋅ ( 40 ⋅ ( 3 ⋅ 10 6 ) ⋅ ( 4 ⋅ 13824 + 2 ⋅ 5120 ) ) ≈ 1.9 GB .

Therefore for memory used per device we get ≈1.9GB.

If we use FSDP and TP to minimize communications we need to set

𝑋 opt = 𝐵 𝐹 ⋅ 𝑀 𝑋 𝑀 𝑌 ⋅ 𝑁

Let 𝑁=163, 𝑀𝑋=2, 𝑀𝑌=1, 𝐵=3⋅106, and 𝐹=13824. Therefore:

𝑋 opt = 3 ⋅ 10 6 13824 ⋅ 2 ⋅ 16 3 ≈ 1333.3 ≈ 1333 .

Since 𝑁=𝑋opt⋅𝑌⇒1631333=𝑌=3.07. To get integer values therefore we choose 𝑋opt=1024 and 𝑌=4, which is closer to our optimum. Having chosen this, we will be compute bound iff:

𝐵 𝑁 > 𝛼 2 𝑀 𝑋 𝑀 𝑌 𝐹

where 𝛼=𝐶𝑊ICI=2550 for v5p.

Let 𝐹=13824, 𝑁=163, 𝑀𝑋=2, 𝑀𝑌=1, and 𝐵=3⋅106:

𝐵 𝑁 = 3 ⋅ 10 6 16 3 ≈ 733 𝛼 2 𝑀 𝑋 𝑀 𝑌 ⋅ 𝐹 = ( 2550 ) 2 2 ⋅ 13824 = 235 .

therefore we are compute bound.

Assuming we are compute bound, using only roofline FLOPs estimates and ignoring attention we get that:

Time per Step = 3 𝑒 6 ⋅ 13 𝑒 9 ⋅ 6 16 3 ⋅ ( 0.4 ⋅ 4.59 𝑒 14 ) = 0.312 s = 312 ms .

Transcription uncertainties

  • The final subpart is labeled 3.2 again in the scan; I preserved the repeated handwritten label.
  • On scan page 6, the slice size appears to be written as $|X| = 4046$; I marked it as uncertain in the transcription.