Chapter 3: Sharded Matrices and How to Multiply Them
Scaling Book Exercises – Chapter 3
Chapter: 3 – Sharded Matrices and How to Multiply Them
Source: handwritten Onyx Boox A4 PDF
Note: Transcribed from handwritten solutions; book markdown used only for exercise statements and notation.
Exercise 1 – replicated sharding
Exercise statement
An array is sharded (i.e., only sharded across ), with a mesh Mesh({'X': 4, 'Y': 8, 'Z': 2}). What is the ratio of the total number of bytes taken up by across all chips to the size of one copy of the array?
Solution
Scan pages: 1-2.
Let and .
At each device tuple , size of array shard is:
Assuming “one copy of array” means unsharded array, ratio is then:
If copy means one shard, then ratio is:
Exercise 2 – AllGather latency
Exercise statement
How long should take on a TPU v4p 4x4x4 slice with mesh Mesh({'X': 4, 'Y': 4, 'Z': 4}) if and in bfloat16? How about ? How about ?
Solution
Scan pages: 2-5.
- HBM capacity =
- HBM BW =
- FLOPs/s (bf16) =
- ICI Bidi =
Slice of . . Let and in bfloat16. Assuming , each device tuple in the mesh maps to:
Therefore size per shard is:
since .
Since v4p slice has wrap-around links, we get that:
For allgather on 2 axis we don’t have an specific algorithm but we can derive a lower bound:
Exercise 3 – latency-bound AllGather
Exercise statement
Let’s say we’re performing an but is very small, say 128. How long should this take on a TPU v4p 4x4x4 slice with mesh Mesh({'X': 4, 'Y': 4, 'Z': 4}) in bfloat16? Hint: you’re probably latency bound.
Solution
Scan pages: 6-7.
TPU v4p.
Let Let
Assuming communication model where , and wraparound links,
So yeah I’m bound by latency.
Exercise 4 – matmul strategies
Exercise statement
To perform , in this section we tell you to perform and multiply the fully replicated matrices (Case 2, Strategy 1). Instead, you could multiply the local shards like (Case 3, Strategy 2), and then . How many FLOPs and comms does each of these perform? Which is better and why?
Solution
Scan pages: 7-16.
Strategy 1:
Note that number of directed links. Each link participates in steps of the bidirectional ring algo and each payload is , assuming bfloat16:
Comms for MatMul
Strategy 2:
From the book they say you can overlap collective communications with the matrix multiplication itself. So then we can work only with the communications of the collectives to see which strategy is better.
So then assume:
Let’s fix valid in v4p slice, and they will still have a wraparound link.
Achievable FLOPs:
Algo 1 branches by .Algo 2 branches by . So in total 4 combinations:
and Both algos are compute bound.
Therefore and in fact it holds for every . Just the condition for when Algo 2 is compute bound changes, but assuming both compute bound strategy 2 is better.
and Algo 1 compute bound and Algo 2 communication bound.
So strategy 1 in this case wins for otherwise strategy 2 is better.
and Algo 1 communication bound and Algo 2 compute bound.
which is a contradiction, so then strategy 2 in this setting is always better, and in fact it will hold for all . It just changes the condition when algo 2 is compute bound.
and Both algos are communication bound.
So if , strategy 1 wins, otherwise strategy 2 wins. In the book they often assume , so then and strategy 2 is better.
Exercise 5 – minimum latency
Exercise statement
Let’s say I want to do a matmul on a TPU v4p 4x4x4 with the lowest possible latency. Assume the inputs can be sharded arbitrarily but the result should be fully replicated. How should my inputs be sharded? What is the total FLOPs and comms time?
Solution
Scan pages: 17-25.
Assume we have a slice of v4p devices with wraparound links. Assume the intended operation is:
with bfloat16 precision.
Note that there is only one operation where we do FLOPs, the matrix multiplication (ignoring AllReduce). This allows to divide the space of possible sharding strategies in three groups:
- Shardings resulting in collective communications before doing matmul.
- Shardings resulting in collective communications only after doing matmul.
- Shardings resulting in collective communications before and after matmul.
In all groups, note that if we are going to do a collective operation, we should choose shardings for which the lower bound in is the smallest, since an optimal algo can get closer to this lower bound. Note that for either we have:
So the smallest lower bound is achieved using shardings that use all the axis of our v4p slice. This argument wipes out our space of possible shardings by a lot.
For group 1:
1.1
1.2
For group 2:
2.1
2.2
2.3
For group 3:
3.1
3.2
3.3
Note that in group 3, depending in which axis you allgather first, you end back to a sharding that belongs to group 1 or 2.
Assume and .
Then for group 1:
1.1 . If
1.2 . If .
Then for group 2:
2.1 . If i.e not realistic as will not fit in HBM.
2.2 . If i.e. . So again not realistic.
2.3 . If i.e . Not realistic.
Hence:
So then Algo 2.2 and 2.1 are better than Algo 1.2 and Algo 1.1 when all are communication bound. And if Algo 1.2 and Algo 1.1 are compute bound, Algo 2.1 and Algo 2.2 are better when:
In that case, between Algo 2.1 and Algo 2.2, Algo 2.2 is better since and:
is dominated by .
is dominated by . So if both memory bound, Algo 2.2 is better. And we enter compute bound with smaller values of with Algo 2.2, since if we are compute bound and for Algo 2.1 we have:
Exercise 6 – sharded matmul cases on TPU v5e 4x4
Exercise statement
Let’s say we want to perform on TPU v5e 4x4. What communication do we perform? How much time is spent on communication vs. computation?
What about ? This is the most standard setting for training where we combine data, tensor, and ZeRO sharding.
What about ? This is standard for inference, where we do pure tensor parallelism plus data.
Solution
Scan pages: 25-28.
Assume we have a slice of TPU v5e. Since is not a full pod in this setting we don’t have wraparound links.
Let:
- v5e ICI BW unidirectional =
- v5e FLOPs/s (bf16) =
Assuming BFLOAT16.
For :
We cannot use a bidirectional ring algorithm for the ReduceScatter + AllGather, since no wraparound links in a slice. Hence:
Therefore:
For :
Therefore:
For :
No communication needs to be done.
All the time is in computation.
Exercise 7 – Transformer block sharding with memory limit
Exercise statement
A typical Transformer block has two matrices and where . Say we have a batch size B. Then the full block is . Let’s pick , , and and assume everything is in bfloat16. Assume we’re running on a TPU v5e 2x2 slice but let’s pretend each TPU only has 300MB of free memory. How should In, , , and Out be sharded to stay below the memory limit while minimizing overall time? How much time is spent on comms and FLOPs? Hint: the final output doesn’t need to be fully replicated, but it should be sharded the same as the input so the layer can be repeated.
Solution
Scan pages: 29-32.
Assume bfloat16. Let , , . The intended operation is:
.
.
.
.
Since we only have left 300MB, then we have to shard and so that both sharded and can be found in each device. Sharding just one dimension along one axis does not work since then both and do not fit in a single device:
Therefore we have to either shard each dimension with a single axis i.e:
where are the axis of the device mesh, or shard one dimension along both axis of the mesh i.e:
In both cases we shrink both and by 4 and:
The way to shard to minimize overall time is:
where we construct a 1-D ring in slice of v5e TPUs.
For our chosen sharding we only do one collective communication at the end, specifically an AllReduce = ReduceScatter + AllGather. Every other sharding strategy will require:
You could also do:
with an AllGather at the beginning and a ReduceScatter at the end and get same:
So for our chosen sharding strategy:
Therefore:
Exercise 9 – another strategy for sharded matmuls?
Exercise statement
Above, we claimed that when only one input to a matmul is sharded along its contracting dimension, we should AllGather the sharded matrix and perform the resulting contraction locally. Another strategy you might think of is to perform the sharded matmul and then AllReduce the result, as if both inputs were sharded along the contracting dimension, i.e. by way of:
Answer the following:
- Explicitly write out this algorithm for matrices and , using indices to show exactly what computation is done on what device. Assume is sharded as across ND devices, and you want your output to be replicated across all devices.
- Now suppose you are ok with the final result not being replicated on each device, but instead sharded, across either the N or K dimension. How would the algorithm above change?
- Looking purely at the communication cost of the strategy above, in part 2 not 1, how does this communication cost compare to the communication cost of the algorithm in which we first AllGather A and then do the matmul?
Solution
Scan pages: 33-34.
9.1
Let be sharded as where . Let be sharded as . Assume .
Define:
for , where we use Python slicing notation and semantics.
Define:
for .
The computation at each device is then . After that we do with an AllReduce.
9.2
9.3
Assume bfloat16.
.
.
.
Exercise 10 – Fun with AllToAll
Exercise statement
In the table above, it was noted that the time to perform an AllToAll is a factor of 4 lower than the time to perform an AllGather or ReduceScatter, in the regime where we are throughput-bound. In this problem we will see where that factor of 4 comes from, and also see how this factor would change if we only had single-direction ICI links, rather than bidirectional ICI links.
- Let’s start with the single-direction case first. Imagine we have D devices in a ring topology and want to do either an AllGather or a ReduceScatter on an N x N matrix , say divides for simplicity. Describe the comms involved in these two collectives, and calculate the total number of scalars, floats or ints, which are transferred across a single ICI link during the entirety of this algorithm.
- Now let’s think about an AllToAll, still in the single-directional ICI case. How is the algorithm different in this case than the all-gather case? Calculate the number of scalars that are transferred across a single ICI link in this algorithm.
- You should have found that the ratio between your answers to part (a) and part (b) is a nice number. Explain where this factor comes from in simple terms.
- Now let’s add bidirectional communication. How does this affect the total time needed in the all-gather case?
- How does adding bidirectional communication affect the total time needed in the AllToAll case?
- Now simply explain the ratio between AllGather time and AllToAll time in a bidirectional ring.
Solution
Scan pages: 35-40.
10.1
Let be sharded as where and .
With only one direction, the ring algo requires steps to finish. Define a link as a directed edge for . Note that at one step of the algo link transfers scalars. Since there are steps, the scalars being transferred by a single link during the entirety of the algorithm is:
for both AllGather and ReduceScatter.
10.2
At step of the AllToAll algorithm, where , the directed link transfers:
Therefore the scalars over the entirety of the algorithm in a single link is:
10.3
It comes from the fact that the payload per link is smaller at each step.
10.4
With bidirectional communication the algo for AllGather takes to finish, so then:
being transferred by a single link or a link .
It takes half the time.
10.5
With bidirectional communication the algorithm for AllToAll takes to finish, so an upper bound is:
scalars being transferred by a single link or . So it takes of the time of the single direction AllToAll.
10.6