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 |
|---|---|
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 this section.
Solution
Scan pages: 1
From LLAMA2 we know that . Therefore:
In total we have parameters. Let , , , , and . Hence:
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
Optimizer State:
Activation Checkpointing:
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.
- Can we use pure data parallelism? Why or why not?
- 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).
- 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:
3.1
We cannot use data parallelism since a lower bound in memory per chip we need is . (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:
and for input output embeddings as:
Therefore the total params per device is:
For LLAMA2 13B therefore:
Assuming optimizer states also sharded, we get that each device needs at least:
If we also take into account the activation checkpointings which are of shapes:
we get:
Let , then:
So memory used per device is .
We could use FSDP since with FSDP sharding everything fits in memory, but we would be communication bound since we are compute bound iff:
Let and , therefore:
and in our slice .
3.3
Assume we shard the attention matrices as:
and for the MLP params as:
and for input and output embeddings as:
Therefore total params per device is:
For LLAMA2 13B we get:
If we take into account activation state then we get:
Activation checkpoints will be sharded with shapes:
so we get:
Let , then:
Therefore for memory used per device we get .
If we use FSDP and TP to minimize communications we need to set
Let , , , , and . Therefore:
Since To get integer values therefore we choose and , which is closer to our optimum. Having chosen this, we will be compute bound iff:
where for v5p.
Let , , , , and :
therefore we are compute bound.
Assuming we are compute bound, using only roofline FLOPs estimates and ignoring attention we get that:
Transcription uncertainties
- The final subpart is labeled
3.2again 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.