How do you split a matrix multiply across GPUs?

Part 2 ended at a wall: the activations of a sequence do not shard along the batch axis, no matter which ZeRO stage you run. To divide them you have to cut the model itself, and the smallest thing worth cutting is one matrix multiply.

There are exactly two ways to cut a weight matrix, and which one you choose is forced by the nonlinearity that follows it, not by the bandwidth. That single observation is most of Megatron’s tensor parallelism, and it is the thing to be able to derive rather than recall.

The rest of the part is arithmetic: how many collectives per layer, how many bytes each one moves, what sequence parallelism adds for free, and why the degree stops at the number of GPUs in a node.

They will ask How does tensor parallelism split a transformer block, and what does it cost?
The 30-second version The first matrix of a pair is split along its columns, so each GPU owns whole output features and any elementwise nonlinearity applies locally with no communication. The second is split along its rows, so it consumes those features directly and produces a partial sum. One all-reduce turns the partial sums back into the residual stream. Attention works the same way with heads as the unit: query, key and value projections are column-parallel, the output projection is row-parallel, and every head's computation stays on one device. Two operators do the bookkeeping and are conjugates: f is identity forward and all-reduce backward, g is all-reduce forward and identity backward. That is two all-reduces per layer in the forward pass and two in the backward, each of sequence times batch times hidden elements. For Llama 3 70B at 8192 tokens and eight-way splitting, that is 134 MB a message and 75 GB per GPU per micro-batch, which is 38% of the arithmetic on NVLink and three and a half times the arithmetic on a 400 Gb/s network. That ratio is why the degree stays inside a node.

Which way do you cut the weight matrix?

Take one GEMM followed by a nonlinearity, $Y = f(XA)$, and two GPUs.

Cut $A$ along its rows and you must cut $X$ along its columns to match. Each GPU computes a partial product, and $f$ cannot be applied to a partial product, because $f(u + v) \neq f(u) + f(v)$ for anything worth calling a nonlinearity. So the partials have to be summed first, which means a synchronisation in the middle of the block.

Cut $A$ along its columns instead and each GPU owns whole output features:

\[[Y_1, Y_2] = [f(XA_1),\; f(XA_2)]\]

The nonlinearity applies column by column, locally, and nothing is communicated at all. That is the asymmetry the whole scheme is built on.

Megatron pairs them. First GEMM column-parallel, second GEMM row-parallel so it consumes the sharded output directly, and one reduction at the end of the pair. The reduction sits exactly where the residual stream has to be whole anyway.

What does that look like in the MLP block?

Llama’s MLP is three matrices, not two: a gate projection and an up projection, both $8192 \times 28672$, and a down projection of $28672 \times 8192$. SwiGLU multiplies the gated half by the up half elementwise.

Split gate and up along their columns. At $t = 8$ each GPU owns 3584 of the 28,672 intermediate features, computes both halves for those features, and applies SwiGLU to its own slice with no communication, because the product is elementwise in the feature index.

Split down along its rows over the same partition. Each GPU multiplies its 3584 features by its 3584 rows and produces a full-width partial sum of the output. One all-reduce adds the eight partial sums.

Each GPU holds 176 MB of this block instead of 1.4 GB, and every activation inside the parallel region is a factor of eight smaller.

What are f and g, exactly?

Two autograd hooks, and they are conjugates of each other. This is the detail that tells an interviewer whether you have implemented it or only read about it.

$f$ sits at the entry to the parallel region. In the forward pass it does nothing, because the input was already broadcast to every GPU in the group. In the backward pass it all-reduces the gradient, because every GPU produced a gradient with respect to that same shared input and they have to be summed.

$g$ sits at the exit. In the forward pass it all-reduces, because the output was a partial sum. In the backward pass it does nothing, because the incoming gradient is already identical on every GPU.

Identity forward, all-reduce backward. All-reduce forward, identity backward. Two lines of PyTorch each.

How does attention split?

More naturally than the MLP, because the heads are already independent.

Give each GPU $64/t$ query heads and the key-value heads that go with them. The query, key and value projections are column-parallel over the head dimension, so each GPU makes only its own heads. The scores, the softmax and the weighted sum for a head happen entirely on the device that owns it: no head ever needs another head.

The output projection is row-parallel over the same partition, so it consumes the local head outputs and emits a partial sum. Same shape as the MLP: column, row, one all-reduce.

There is a hard limit hiding in there, and it is grouped-query attention. Llama 3 70B has 8 key-value heads. At $t = 8$ each GPU owns exactly one, which is tidy. At $t = 16$ the key-value heads have to be duplicated across pairs of GPUs, and duplicated weights mean duplicated gradients and an extra reduction to keep them in step. The KV head count is a ceiling on the tensor-parallel degree that has nothing to do with bandwidth.

How many collectives is that, and how many bytes?

Two all-reduces per layer in the forward pass, one after attention and one after the MLP. Two more in the backward pass from the conjugate operators. Four per layer, and for 80 layers that is 320 collectives per micro-batch.

Each message is $s\,b\,h$ elements. At 8192 tokens and $h = 8192$ in BF16, 134.2 MB. A ring all-reduce over $t$ ranks moves $2(t-1)/t$ of that per rank, which at $t = 8$ is 234.9 MB.

\[4 \times 80 \times 234.9\ \text{MB} = 75.2\ \text{GB per GPU per micro-batch}\]

Against 438 ms of arithmetic on this GPU, that is 167 ms on NVLink at 450 GB/s per direction, so 38% overhead. On a 400 Gb/s network port at 50 GB/s it would be 1.50 seconds, three and a half times the arithmetic.

And note what it is per. Data parallelism’s collective fires once per optimizer step and amortises over gradient accumulation. This one fires per micro-batch, four times per layer, and amortises over nothing.

Key idea Tensor parallelism is the only axis that divides both the weights and the activations, and it pays for that with the highest-frequency collective in the job. Its degree is chosen by the interconnect topology, not by the model.

What is sequence parallelism, and why is it free?

Tensor parallelism leaves two things replicated on every GPU in the group: the layer norms and the dropouts, and the residual stream they act on. They are cheap to compute and expensive to store.

In Megatron’s activation accounting they are the $10\,s\,b\,h$ that stubbornly refuses to divide:

\[\text{tensor parallel: } s\,b\,h\left(10 + \frac{24}{t}\right) \qquad \text{tensor + sequence: } s\,b\,h\left(\frac{34}{t}\right)\]

The fix is that those operations are independent along the sequence, so they can be split that way instead. The all-reduce at the end of a parallel region becomes a reduce-scatter, which lands each GPU with its own slice of the sequence, and the entry to the next region becomes an all-gather.

The bytes on the wire do not change, because a ring all-reduce was a reduce-scatter followed by an all-gather all along. Korthikanti and colleagues make that point explicitly: four all-reduces per forward and backward becomes four all-gathers and four reduce-scatters, at identical bandwidth.

What changes is memory. At $t = 8$ and our shape, a layer goes from 872 MB to 285 MB, a factor of three, and the model’s activations from 69.8 GB to 22.8 GB. It is the rare optimisation with no downside, which is why every framework bundles it with tensor parallelism rather than exposing it as a separate switch.

What does tensor parallelism do to each kernel?

Two things, and they both get worse as $t$ grows.

The matmuls get narrower. The MLP GEMM on one GPU is $8192 \times 3584$ at $t = 8$, and $8192 \times 896$ at $t = 32$. A narrow GEMM has fewer tiles to spread over 132 streaming multiprocessors and a worse ratio of epilogue to arithmetic, so it sits further below the roofline. This is the Part 2 of the GPU series problem arriving through a different door.

And every all-reduce is a synchronisation. Whatever the fastest GPU in the group finishes early is wasted, four times per layer, so a straggler tax that would be invisible once per step becomes 320 taxes per micro-batch.

There is a live line of work on the first problem. Asynchronous tensor parallelism decomposes the GEMM and the collective into chunks so that a slice of the matmul overlaps the transfer of the previous slice, using symmetric memory to write directly into a peer’s buffer. It is in torchtitan under compile, and it is the main reason the practical ceiling on $t$ has moved at all since 2021.

Why does the degree stop at eight?

Three limits arrive at almost the same place, which is unusual and convenient.

Bandwidth. Arithmetic per GPU falls as $1/t$, while communication climbs towards a ceiling, because $2(t-1)/t \to 2$. So the ratio of communication to computation grows roughly as $t$: 5% at $t=2$, 16% at $t=4$, 38% at $t=8$, 82% at $t=16$ on NVLink. Across a network it is nine times worse and hopeless everywhere.

Heads. Eight key-value heads means eight is the largest degree that does not duplicate them.

Tiles. At $t = 32$ the FFN matmul is 896 wide, which is a handful of tile columns.

Megatron’s own guidance from the other direction is the same conclusion: use tensor parallelism up to the number of GPUs in a server, then use pipeline parallelism to go further. Every published dense recipe in Part 6 has $t = 8$, and the reason is that servers have eight GPUs.

They will ask You need to halve the activation memory. Do you raise the tensor-parallel degree from 8 to 16?
The 30-second version Almost certainly not, and the reasons are worth listing in order. The group would straddle two nodes, so 320 collectives per micro-batch would move from a 450 GB/s link to a 50 GB/s one, and the arithmetic they overlap with has halved at the same time, so the ratio goes from about 38% to well over 100%. The model has eight key-value heads, so sixteen-way splitting duplicates them and adds a reduction. And the FFN matmul narrows to 1792 columns, which costs efficiency inside the kernel. Cheaper options first: turn on sequence parallelism if it is somehow off, since it is a factor of three on this shape for no extra bytes; enable selective activation recomputation, which drops the attention internals for a few percent of extra compute; halve the micro-batch and use gradient accumulation to keep the global batch; or split the sequence with context parallelism if the sequence is what is long. Raising the tensor degree past the node is the last thing on the list, not the first.

Rapid fire: can you do these from memory?

  1. Explain why the first GEMM of a pair is column-parallel and the second is row-parallel.
  2. Define f and g and say what each does in each pass.
  3. Count the collectives per layer per micro-batch, and give the message size in elements.
  4. Compute the per-GPU tensor-parallel traffic for 80 layers at 8192 tokens and eight-way splitting.
  5. What does sequence parallelism split, what does it cost, and which term of the activation formula does it fix?
  6. Why does grouped-query attention cap the tensor-parallel degree?
  7. Name three separate reasons the degree stops at the size of a node.

Part 4 takes the other way of cutting the model: not across each layer, but between them.