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 Mesh={𝑋:4,𝑌:8,𝑍:2}.

At each device tuple (𝑖,𝑗,𝑘), size of array shard is:

Size of shard = | 𝐼 | 𝑋 ⋅ 𝑃 ⋅ BytesForPrecision where 𝑃 = ∏ 𝑑 ∈ dimensions − { 𝐼 } | 𝑑 |

Assuming “one copy of array” means unsharded array, ratio is then:

𝑋 ⋅ 𝑌 ⋅ 𝑍 ⋅ | 𝐼 | ⋅ 𝑋 − 1 ⋅ 𝑃 ⋅ BytesForPrecision | 𝐼 | ⋅ 𝑃 ⋅ BytesForPrecision = 16

If copy means one shard, then ratio is:

𝑋 ⋅ 𝑌 ⋅ 𝑍 ⋅ | 𝐼 | ⋅ 𝑋 − 1 ⋅ 𝑃 ⋅ BytesForPrecision | 𝐼 | ⋅ 𝑋 − 1 ⋅ 𝑃 ⋅ BytesForPrecision = 64

Exercise 2 – AllGather latency

Exercise statement

How long should AllGather𝑋([𝐵𝑋,𝐷𝑌]) take on a TPU v4p 4x4x4 slice with mesh Mesh({'X': 4, 'Y': 4, 'Z': 4}) if 𝐵=1024 and 𝐷=4096 in bfloat16? How about AllGather𝑋𝑌([𝐵𝑋,𝐷𝑌])? How about AllReduce𝑍([𝐵𝑋,𝐷𝑌]{𝑈𝑍})?

Solution

Scan pages: 2-5.

v4p:

  • HBM capacity = 32GB
  • HBM BW = 1.2𝑒12Bytes/s
  • FLOPs/s (bf16) = 2.75𝑒14
  • ICI Bidi = 9.0𝑒10

Slice of 4×4×4 v4p. Mesh={𝑋:4,𝑌:4,𝑍:4}. Let 𝐵=1024 and 𝐷=4096 in bfloat16. Assuming 𝐴[𝐵𝑋,𝐷𝑌], each device tuple in the mesh maps to:

( 𝑖 , 𝑗 , 𝑘 ) ↦ 𝐴 [ 𝑖 ⋅ | 𝐵 | 𝑋 : ( 𝑖 + 1 ) ⋅ | 𝐵 | 𝑋 , 𝑗 ⋅ | 𝐷 | 𝑌 : ( 𝑗 + 1 ) ⋅ | 𝐷 | 𝑌 ]

Therefore size per shard is:

| 𝐵 | 𝑋 ⋅ | 𝐷 | 𝑌 ⋅ 2 = 𝐵 2 2 bytes

since 𝐷=4⋅𝐵.

Since v4p slice 4×4×4 has wrap-around links, we get that:

𝑇 ( AllGather 𝑋 ( [ 𝐵 𝑋 , 𝐷 𝑌 ] ) ) = 𝑋 2 ⋅ 𝐵 2 2 ⋅ 𝑊 uni = 𝐵 2 4.5 𝑒 10 ≈ 2 ⋅ 10 − 5 s = 0.02 ms.

For allgather on 2 axis we don’t have an specific algorithm but we can derive a lower bound:

𝑇 ( AllGather 𝑋 𝑌 ( [ 𝐵 𝑋 , 𝐷 𝑌 ] ) ) ≥ 2 ⋅ 𝐵 ⋅ 𝐷 2 ⋅ 𝑊 ICI = 2 𝐵 ⋅ 4 𝐵 2 𝑊 ICI = 4 𝐵 2 9 𝑒 10 ≈ 0.04 ms.

𝑇(AllReduce𝑍([𝐵𝑋,𝐷𝑌]{𝑈𝑍}))=𝑍2⋅(𝐵22⋅𝑍⋅𝑊uni)+𝑍2⋅(𝐵22⋅𝑍⋅𝑊uni)=𝐵24⋅𝑊uni+𝐵24⋅𝑊uni=2𝐵24𝑊uni=𝐵22𝑊uni=𝐵2𝑊ICI=102429𝑒10=1⋅10−4=0.01ms.

Exercise 3 – latency-bound AllGather

Exercise statement

Let’s say we’re performing an AllGather𝑋([𝐵𝑋]) 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 Mesh={𝑋:4,𝑌:4,𝑍:4} Let 𝐵=128and𝐴=bfloat16[𝐵𝑋]

Assuming communication model 𝑇step=𝛼+𝐷BW where 𝛼=1𝜇s, and wraparound links, 𝑇(AllGather𝑋([𝐵𝑋]))=𝑋2⋅(𝛼+𝐵⋅2𝑋⋅𝑊uni)=2⋅𝛼+2⋅𝐵𝑊ICI=2𝜇s+2569𝑒10≈2𝜇s+28.5⋅10−10=2𝜇s+28.5⋅10−3⋅10−6=2𝜇s+0.0285𝜇s

So yeah I’m bound by latency.

Exercise 4 – matmul strategies

Exercise statement

To perform 𝑋[𝐵,𝐷]⋅𝐷𝑌[𝐷𝑋,𝐹]→𝑍[𝐵,𝐹], in this section we tell you to perform AllGather𝑋(𝑌[𝐷𝑋,𝐹]) and multiply the fully replicated matrices (Case 2, Strategy 1). Instead, you could multiply the local shards like 𝑋[𝐵,𝐷𝑋]⋅𝐷𝑌[𝐷𝑋,𝐹]→𝑍[𝐵,𝐹]{𝑈𝑋} (Case 3, Strategy 2), and then AllReduce𝑋(𝑍[𝐵,𝐹]{𝑈𝑋}). How many FLOPs and comms does each of these perform? Which is better and why?

Solution

Scan pages: 7-16.

𝑋 [ 𝐵 , 𝐷 ] ⋅ 𝐷 𝑌 [ 𝐷 𝑋 , 𝐹 ] → 𝑍 [ 𝐵 , 𝐹 ]

Strategy 1:

  1. AllGather𝑋(𝑌[𝐷𝑋,𝐹])→𝑌[𝐷,𝐹]
  2. 𝑋[𝐵,𝐷]⋅𝐷𝑌[𝐷,𝐹]→𝑍[𝐵,𝐹]

Note that 2⋅𝑋= number of directed links. Each link participates in 𝑋2 steps of the bidirectional ring algo and each payload is 𝐷⋅𝐹⋅2𝑋, assuming bfloat16:

AllGather [ 𝐷 𝑋 , 𝐹 ] Comms = 2 ⋅ ( 𝑋 2 ⋅ 𝐷 ⋅ 𝐹 ⋅ 2 𝑋 ) = 2 𝐷 𝐹

Comms for MatMul =2𝐵𝐷+2𝐷𝐹+2𝐵𝐹 ⇒Total comms strategy 1=2𝐷𝐹+2𝐵𝐷+2𝐷𝐹+2𝐵𝐹=4𝐷𝐹+2𝐵𝐷+2𝐵𝐹.

FLOPs for strategy 1=2⋅𝐵⋅𝐷⋅𝐹

Strategy 2:

  1. 𝑋[𝐵,𝐷𝑋]⋅𝐷𝑌[𝐷𝑋,𝐹]→𝑍[𝐵,𝐹]{𝑈𝑋}
  2. AllReduce𝑋(𝑍[𝐵,𝐹]{𝑈𝑋})

Comms MatMul=2⋅𝐵⋅𝐷𝑋+2⋅𝐷𝑋⋅𝐹+2⋅𝐵⋅𝐹.

CommsAllReduce𝑋(𝑍[𝐵,𝐹]{𝑈𝑋})=2⋅(2(𝑋2⋅𝐵⋅𝐹⋅2𝑋))=4𝐵𝐹.

Total comms strategy 2=4𝐵𝐹+2⋅𝐵⋅𝐷𝑋+2⋅𝐷𝑋⋅𝐹+2𝐵𝐹=6𝐵𝐹+2⋅𝐵⋅𝐷𝑋+2⋅𝐷𝑋⋅𝐹.

FLOPs strategy 2=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:

𝑇 strategy1 = max ( 2 𝐷 𝐹 𝑊 ICI , 2 𝐵 𝐷 𝐹 FLOPs/s ) 𝑇 strategy2 = max ( 4 𝐵 𝐹 𝑊 ICI , 2 𝐵 𝐷 𝐹 𝑋 ⋅ FLOPs/s ) 𝐴 𝐼 1 = 2 𝐵 𝐷 𝐹 2 𝐷 𝐹 = 𝐵 FLOPs/Byte 𝐴 𝐼 2 = 2 𝐵 𝐷 𝐹 ⋅ 𝑋 − 1 4 𝐵 𝐹 = 𝐷 2 𝑋 FLOPs/Byte

Let’s fix 𝑋∈{4,8,16} valid in v4p slice, and they will still have a wraparound link. 𝐴𝐼1(𝐵),𝐴𝐼2(𝐷;𝑋)for fixed values of𝑋.

Achievable FLOPs:

Achievable FLOPs ( 𝐵 ) = min ( 𝐵 ⋅ 9 𝑒 10 , 2.75 𝑒 14 ) Achievable FLOPs ( 𝐷 ; 𝑋 ) = min ( 𝐷 2 𝑋 ⋅ 9 𝑒 10 , 2.75 𝑒 14 )

Algo 1 branches by 𝐵>3055.Algo 2 branches by 𝐷>24𝐾. So in total 4 combinations:

𝐵>3055 and 𝐷>24𝐾 ⇒ Both algos are compute bound.

𝑇 math ( Strat 1 ) < 𝑇 math ( Strat 2 ) ⇔ 2 𝐵 𝐷 𝐹 < 2 𝐵 𝐷 𝐹 𝑋 for 𝑋 = 4 ⇔ 2 < 1 2 Contradiction.

Therefore 𝑇math(Strat 2)≥𝑇math(Strat 1) and in fact it holds for every 𝑋>1. Just the condition for when Algo 2 is compute bound changes, but assuming both compute bound strategy 2 is better.

𝐵>3055 and 𝐷<24𝐾 ⇒ Algo 1 compute bound and Algo 2 communication bound.

𝑇 math ( Strat 1 ) < 𝑇 comms ( Strat 2 ) ⇔ 2 𝐵 𝐷 𝐹 FLOPs/s < 4 𝐵 𝐹 𝑊 ICI ⇔ 2 𝐷 4 < FLOPs 𝑊 ICI = 3055 ⇔ 𝐷 2 < 3055 ⇒ 𝐷 < 6110

So strategy 1 in this case wins for 𝐷<6110 otherwise strategy 2 is better.

𝐵<3055 and 𝐷>24𝐾 ⇒ Algo 1 communication bound and Algo 2 compute bound.

𝑇 comms ( Strat 1 ) < 𝑇 math ( Strat 2 ) ⇔ 2 𝐷 𝐹 𝑊 ICI < 2 𝐵 𝐷 𝐹 𝑋 ⋅ FLOPs/s ) for 𝑋 = 4 ⇔ 2 𝑊 ICI < 2 𝐵 4 ⋅ FLOPs/s ⇔ 𝐵 > 4 ⋅ FLOPs/s 𝑊 ICI = 4 ⋅ 3055 ≈ 12 𝐾

which is a contradiction, so then strategy 2 in this setting is always better, and in fact it will hold for all 𝑋>2. It just changes the condition when algo 2 is compute bound.

𝐵≤3055 and 𝐷≤24𝐾 ⇒ Both algos are communication bound.

𝑇 comms ( Strat 1 ) < 𝑇 comms ( Strat 2 ) ⇔ 2 𝐷 𝐹 < 4 𝐵 𝐹 ⇔ 2 𝐷 < 4 𝐵 ⇒ 𝐷 < 2 ⋅ 𝐵

So if 𝐷<2⋅𝐵, strategy 1 wins, otherwise strategy 2 wins. In the book they often assume 𝐷≫𝐵, so then 𝐷>2⋅𝐵 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 4×4×4 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:

  1. Shardings resulting in collective communications before doing matmul.
  2. Shardings resulting in collective communications only after doing matmul.
  3. 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 𝑇comms is the smallest, since an optimal algo can get closer to this lower bound. Note that for either 𝑂∈{𝐴,𝐵,𝐶} we have:

𝑇 comms ≥ size ( 𝑂 ) 𝑁 axes ⋅ 𝑊 ICI .

So the smallest lower bound is achieved using shardings that use all the axis of our 4×4×4 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 𝐾≈4⋅𝐽.

Then for group 1:

1.1 𝑇comms=2⋅𝐼⋅𝐽,𝑇math=2⋅𝐼⋅𝐽⋅𝐾. If 2𝐼𝐽𝐾2𝐼𝐽=𝐾>3055⇒compute bound

1.2 𝑇comms=2⋅𝐽⋅𝐾,𝑇math=2⋅𝐼⋅𝐽⋅𝐾. If 𝐼>3055⇒compute bound.

Then for group 2:

2.1 𝑇comms=2⋅𝐼⋅𝐾,𝑇math=2𝐼𝐽𝐾𝑋⋅𝑌⋅𝑍. If 𝐽𝑋⋅𝑌⋅𝑍>3055⇒compute bound i.e 𝐽>64⋅3055=196𝐾 not realistic as 𝐵[𝐽,𝐾] will not fit in HBM.

2.2 𝑇comms=2⋅𝐼⋅𝐾,𝑇math=2𝐼𝐽𝐾𝑋⋅𝑌⋅𝑍. If 𝐽𝑋⋅𝑌⋅𝑍>3055⇒compute bound i.e. 𝐽>196𝐾. So again not realistic.

2.3 𝑇comms=4⋅𝐼⋅𝐾,𝑇math=2𝐼𝐽𝐾𝑋⋅𝑌⋅𝑍. If 𝐽2𝑋⋅𝑌⋅𝑍>3055⇒compute bound i.e 𝐽>391𝐾. Not realistic.

Hence:

𝑇 ( Algo 1.1 ) = 2 𝐼 𝐽 𝐾 FLOPs/s ≈ 8 𝐽 2 FLOPs/s 𝑇 ( Algo 1.2 ) = 2 𝐼 𝐽 𝐾 FLOPs/s ≈ 8 𝐽 2 FLOPs/s or = 2 𝐽 𝐾 𝑊 ICI = 8 𝐽 2 𝑊 ICI 𝑇 ( Algo 2.2 ) = 2 𝐼 𝐾 𝑊 ICI ≈ 8 𝐽 𝑊 ICI

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:

8 𝐽 𝑊 ICI < 8 𝐽 2 FLOPs/s ⇒ 𝐽 > 3055

In that case, between Algo 2.1 and Algo 2.2, Algo 2.2 is better since 𝐾≫𝐼 and:

comms Algo 2.1 = 2 𝐼 𝐽 𝑋 𝑌 𝑍 + 2 𝐽 𝐾 + 2 𝐼 𝐾

is dominated by 2𝐽𝐾.

comms Algo 2.2 = 2 𝐼 𝐾 + 2 𝐽 𝐾 𝑋 𝑌 𝑍 + 2 𝐼 𝐾

is dominated by 2𝐽𝐾𝑋𝑌𝑍. 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 𝐼=𝐴𝐼(Algo 2.2)>230 we are compute bound and for Algo 2.1 we have:

𝐴 𝐼 = 2 𝐼 𝐽 𝐾 𝑋 𝑌 𝑍 ⋅ 1 2 𝐽 𝐾 = 2 64 ⋅ 𝐼 > 230 ⇒ 𝐼 > 7360

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 4×4 slice of TPU v5e. Since 4×4 is not a full pod in this setting we don’t have wraparound links.

Let:

  • v5e ICI BW unidirectional = 4.5𝑒10
  • v5e FLOPs/s (bf16) = 1.97𝑒14

Assuming BFLOAT16.

For 𝐴[𝐼𝑋,𝐽𝑌]⋅𝐽𝐵[𝐽𝑌,𝐾]→𝐶[𝐼𝑋,𝐾]:

  1. 𝑂[𝐼𝑋,𝐾]{𝑈𝑌}=𝐴[𝐼𝑋,𝐽𝑌]⋅𝐽local[𝐽𝑌,𝐾]
  2. 𝐶[𝐼𝑋,𝐾]=AllReduce𝑌(𝑂[𝐼𝑋,𝐾]{𝑈𝑌})

We cannot use a bidirectional ring algorithm for the ReduceScatter + AllGather, since no wraparound links in a 4×4 slice. Hence:

Communication Load = 2 ⋅ ( | 𝑌 | − 1 ) ⋅ 𝐼 ⋅ 𝐾 ⋅ 2 | 𝑌 | | 𝑋 | ≈ 2 ⋅ 𝐼 ⋅ 𝐾 ⋅ 2 | 𝑋 | ⇒ 𝑇 comms = 4 𝐼 𝐾 | 𝑋 | ⋅ BW 𝑇 math = 2 𝐼 𝐽 𝐾 | 𝑋 | | 𝑌 | ⋅ FLOPs/s .

Therefore:

𝑇 comms 𝑇 math = 4 𝐼 𝐾 | 𝑋 | ⋅ BW ⋅ | 𝑋 | | 𝑌 | ⋅ FLOPs/s 2 𝐼 𝐽 𝐾 = 4 | 𝑌 | 𝐽 ⋅ FLOPs/s BW = 16 𝐽 ⋅ 4380 ≈ 70080 𝐽 .

For 𝐴[𝐼𝑋,𝐽]⋅𝐽𝐵[𝐽𝑋,𝐾𝑌]→𝐶[𝐼𝑋,𝐾𝑌]:

  1. 𝐵[𝐽,𝐾𝑌]=AllGather𝑋(𝐵[𝐽𝑋,𝐾𝑌])
  2. 𝐶[𝐼𝑋,𝐾𝑌]=𝐴[𝐼𝑋,𝐽]⋅𝐵[𝐽,𝐾𝑌]
Communication load = ( | 𝑋 | − 1 ) ⋅ ( 2 𝐽 ⋅ 𝐾 | 𝑋 | | 𝑌 | ) ≈ 2 ⋅ 𝐽 ⋅ 𝐾 ⋅ | 𝑌 | − 1 ⇒ 𝑇 comms = 2 𝐽 𝐾 | 𝑌 | ⋅ BW 𝑇 math = 2 𝐼 𝐽 𝐾 | 𝑋 | | 𝑌 | ⋅ FLOPs/s .

Therefore:

𝑇 comms 𝑇 math = 2 𝐽 𝐾 | 𝑌 | ⋅ | 𝑋 | | 𝑌 | ⋅ FLOPs/s 2 𝐼 𝐽 𝐾 ⋅ BW = | 𝑋 | 𝐼 ⋅ 4380 = 17520 𝐼 .

For 𝐴[𝐼𝑋,𝐽]⋅𝐽𝐵[𝐽,𝐾𝑌]→𝐶[𝐼𝑋,𝐾𝑌]:

No communication needs to be done.

𝑇 math = 2 𝐼 𝐽 𝐾 | 𝑋 | | 𝑌 | ⋅ FLOPs/s

All the time is in computation.

Exercise 7 – Transformer block sharding with memory limit

Exercise statement

A typical Transformer block has two matrices 𝑊in[𝐷,𝐹] and 𝑊out[𝐹,𝐷] where 𝐹≫𝐷. Say we have a batch size B. Then the full block is In[𝐵,𝐷]⋅𝑊in[𝐷,𝐹]⋅𝑊out[𝐹,𝐷]. Let’s pick 𝐷=8192, 𝐹=32768, and 𝐵=128 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, 𝑊in, 𝑊out, 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 𝐷=8192, 𝐹=32768, 𝐵=128. The intended operation is:

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

Size(𝑊in)=81922⋅4⋅2≈536MB.

Size(𝑊out)=Size(𝑊in).

Size(In)≈2MB.

Size(Aux[𝐵,𝐹])≈8MB.

Since we only have left 300MB, then we have to shard 𝑊in and 𝑊out so that both sharded 𝑊in and 𝑊out can be found in each device. Sharding just one dimension along one axis does not work since then both 𝑊in and 𝑊out do not fit in a single device:

Size ( 𝑊 in ) 2 + Size ( 𝑊 out ) 2 > 300 MB

Therefore we have to either shard each dimension with a single axis i.e:

𝑊 in [ 𝐷 𝑋 , 𝐹 𝑌 ] ∀ ∈ { 𝑋 , 𝑌 }

where 𝑋,𝑌 are the axis of the device mesh, or shard one dimension along both axis of the mesh i.e:

𝑊 in [ 𝐷 , 𝐹 𝑋 𝑌 ] or 𝑊 in [ 𝐷 𝑋 𝑌 , 𝐹 ]

In both cases we shrink both Size(𝑊in) and Size(𝑊out) by 4 and:

Size ( 𝑊 in ) 4 + Size ( 𝑊 out ) 4 ≤ 300 MB

The way to shard to minimize overall time is:

In [ 𝐵 , 𝐷 ] @ 𝑊 in [ 𝐷 , 𝐹 𝑋 𝑌 ] @ 𝑊 out [ 𝐹 𝑋 𝑌 , 𝐷 ]

where we construct a 1-D ring in 2×2 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:

𝑇 comms ( Alternative sharding ) ≥ 4 𝐵 𝐷 𝑊 ICI

You could also do:

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

with an AllGather at the beginning and a ReduceScatter at the end and get same:

𝑇 comms = 2 𝐵 𝐷 𝑊 ICI + 2 𝐵 𝐷 𝑊 ICI = 4 𝐵 𝐷 𝑊 ICI

So for our chosen sharding strategy:

  1. Aux[𝐵,𝐹𝑋𝑌]=In[𝐵,𝐷]⋅𝑊in[𝐷,𝐹𝑋𝑌]
  2. 𝑂[𝐵,𝐷]{𝑈𝑋𝑌}=Aux[𝐵,𝐹𝑋𝑌]⋅𝑊out[𝐹𝑋𝑌,𝐷]
  3. 𝑂[𝐵,𝐷]=AllReduce𝑋𝑌(𝑂[𝐵,𝐷]{𝑈𝑋𝑌})

Therefore:

𝑇 math = 4 𝐵 𝐷 𝐹 | 𝑋 | | 𝑌 | ⋅ 1.97 𝑒 14 = 128 ⋅ 8192 ⋅ 32768 1.97 𝑒 14 ≈ 0.0002 s = 0.2 ms 𝑇 comms = 4 𝐵 𝐷 9.0 𝑒 10 ≈ 0.00004 s = 4 ⋅ 10 − 5 = 4 ⋅ 10 − 2 ⋅ 10 − 3 s = 0.04 ms

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:

  1. 𝐶[𝐼,𝐾]{𝑈𝑋}=𝐴[𝐼,𝐽𝑋]⋅𝐵[𝐽𝑋,𝐾]
  2. 𝐶[𝐼,𝐾]=AllReduce(𝐶[𝐼,𝐾]{𝑈𝑋})

Answer the following:

  1. 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.
  2. 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?
  3. 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:

𝐴 𝑖 ≔ 𝐴 [ : , 𝑖 ⋅ 𝑀 𝐷 : ( 𝑖 + 1 ) ⋅ 𝑀 𝐷 ]

for 0≤𝑖≤𝐷−1, where we use Python slicing notation and semantics.

Define:

𝐵 𝑖 ≔ 𝐵 [ 𝑖 ⋅ 𝑀 𝐷 : ( 𝑖 + 1 ) ⋅ 𝑀 𝐷 , : ]

for 0≤𝑖≤𝐷−1.

The computation at each device 𝑖 is then 𝐴𝑖⋅𝐵𝑖∈ℝ𝑁×𝐾. After that we do 𝐶[𝐼,𝐾]=∑𝑖=0𝐷−1𝐴𝑖⋅𝐵𝑖 with an AllReduce.

9.2

  1. 𝐶[𝐼,𝐾]{𝑈𝑋}=𝐴[𝐼,𝐽𝑋]⋅𝐵[𝐽𝑋,𝐾]
  2. ReduceScatter𝑋,𝐼or𝐾(𝐶[𝐼,𝐾]{𝑈𝑋})

9.3

Assume bfloat16.

Communication Load Part 2=2⋅𝑁⋅𝐾.

Communication load AllGather first=2⋅𝑀⋅𝐾.

Ratio=2⋅𝑁⋅𝐾2⋅𝑀⋅𝐾=𝑁𝑀.

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.

  1. 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.
  2. 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.
  3. 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.
  4. Now let’s add bidirectional communication. How does this affect the total time needed in the all-gather case?
  5. How does adding bidirectional communication affect the total time needed in the AllToAll case?
  6. 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 𝐷−1 steps to finish. Define a link 𝑖 as a directed edge (𝑖,𝑖+1mod𝐷) for 0≤𝑖≤𝐷−1. Note that at one step of the algo link 𝑖 transfers 𝑁2𝐷 scalars. Since there are 𝐷−1 steps, the scalars being transferred by a single link during the entirety of the algorithm is:

( 𝐷 − 1 ) ⋅ 𝑁 2 𝐷

for both AllGather and ReduceScatter.

10.2

At step 𝑘 of the AllToAll algorithm, where 1≤𝑘≤𝐷−1, the directed link (𝑖,𝑖+1mod𝐷) transfers:

𝑁 2 𝐷 − 𝑘 ⋅ 𝑁 2 𝐷 2

Therefore the scalars over the entirety of the algorithm in a single link is:

∑ 𝑘 = 1 𝐷 − 1 ( 𝑁 2 𝐷 − 𝑘 ⋅ 𝑁 2 𝐷 2 ) = ( 𝐷 − 1 ) 𝑁 2 𝐷 − ( 𝐷 − 1 ) 𝐷 ⋅ 𝑁 2 2 ⋅ 𝐷 2 = 𝐷 − 1 𝐷 ⋅ ( 𝑁 2 − 𝑁 2 2 )

10.3

Scalars ( AllGather ) = ( 𝐷 − 1 ) ⋅ 𝑁 2 𝐷 Scalars ( AllToAll ) = 𝐷 − 1 𝐷 ⋅ ( 𝑁 2 − 𝑁 2 2 ) Scalars ( AllToAll ) Scalars ( AllGather ) = 1 2

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 𝐷2 to finish, so then:

𝐷 2 ⋅ ( 𝑁 2 𝐷 ) = 𝑁 2 2 scalars

being transferred by a single link (𝑖,𝑖+1mod𝐷) or a link (𝑖+1mod𝐷,𝑖).

Scalars ( AllGather Bidi ) Scalars ( AllGather ) = 𝑁 2 2 ( 𝐷 − 1 ) ⋅ 𝑁 2 𝐷 = 1 2 ⋅ 𝐷 𝐷 − 1 ≈ 1 2 as 𝐷 → ∞

It takes half the time.

10.5

With bidirectional communication the algorithm for AllToAll takes 𝐷2 to finish, so an upper bound is:

∑ 𝑘 = 1 𝐷 2 ( 𝑁 2 𝐷 ⋅ 2 − 𝑘 𝑁 2 𝐷 2 ) ≈ 𝐷 2 ⋅ 𝑁 2 𝐷 ⋅ 2 − 𝑁 2 𝐷 2 ⋅ 𝐷 2 8 = 𝑁 2 4 − 𝑁 2 8 = 1 4 ⋅ ( 𝑁 2 − 𝑁 2 2 )

scalars being transferred by a single link (𝑖,𝑖+1mod𝐷) or (𝑖+1mod𝐷,𝑖). So it takes 14 of the time of the single direction AllToAll.

10.6

Scalars ( Bidi AllToAll ) Scalars ( Bidi AllGather ) = 𝑁 2 8 ⋅ 2 𝑁 2 = 2 8 = 1 4