← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 10 · Chapters 21–22

Tensor Parallelism

Split the matrix, not the model. Pipeline parallelism cuts between layers; tensor parallelism cuts inside a single matrix multiplication. The whole design reduces to one question: after you split a matmul, is the result already correct, or is it a partial sum that must be added up?

Two ways to split one matmul

Take $Y=XA$ with $X\in\mathbb R^{n\times k}$ and $A\in\mathbb R^{k\times d}$. There are exactly two axes of $A$ to cut, and they behave completely differently.

Column parallel

Split $A$ by columns, $A=[A_1\ A_2]$. Each rank holds all of $X$ and one column block:

$$Y=X[A_1\ A_2]=[\,XA_1\ \ XA_2\,]=[\,Y_1\ \ Y_2\,].$$

Rank $i$ computes $Y_i$, a genuine slice of the answer — not a partial sum. No forward communication. The output is left sharded along its last dimension.

Row parallel

Split $B$ by rows and $X$ by columns to match:

$$Y=[\,X_1\ X_2\,]\begin{bmatrix}B_1\\B_2\end{bmatrix}=X_1B_1+X_2B_2.$$

Now rank $i$ produces $X_iB_i$, which is a partial sum of full shape. The pieces must be added: one all-reduce.

Column parallel

Input replicated, output sharded. Forward: nothing. Backward: all-reduce the input gradient, because $X$ was replicated and replication is a fan-out.

Row parallel

Input sharded, output replicated. Forward: all-reduce. Backward: nothing, because each rank contributed one addend.

The conjugate pair. These are the first two rows of the duality table from the autograd note, wrapped as autograd functions. Megatron calls them $f$ and $\bar f$:
$$f:\ \text{identity forward},\ \text{all-reduce backward};\qquad \bar f:\ \text{all-reduce forward},\ \text{identity backward}.$$
Implement those two and every sharded layer differentiates correctly with no further thought.

Why an MLP needs only one all-reduce

Naively you would expect a collective after each of the two matmuls. Chaining the two layouts avoids one of them entirely.

1

Make the up projection column parallel. Input is replicated, output is sharded along the hidden dimension. No communication.

2

Apply the nonlinearity on the shard. This is the crucial step: SiLU and the gating multiply are elementwise, so $\phi$ applied to a slice equals the slice of $\phi$ applied to the whole. No communication, and no approximation.

3

Make the down projection row parallel. Its input is already sharded exactly the way row-parallel wants it. One all-reduce produces the replicated output.

$$\text{MLP}=\underbrace{\bar f}_{\text{one all-reduce}}\Big(W_{\text{down}}^{(i)}\ \phi\big(W_{\text{gate}}^{(i)}x\big)\odot W_{\text{up}}^{(i)}x\Big).$$
An elementwise nonlinearity is what makes this work. If the activation mixed hidden units — a softmax or a normalization across the hidden dimension — the shard would be wrong and a collective would be required in the middle. This is precisely why the normalization layers get special treatment below.

Attention splits along heads

Attention is even more convenient, because it is already parallel: heads never interact until the output projection.

TensorLayoutCommunication
$W_Q,W_K,W_V$column parallel — each rank owns $h/P$ whole headsnone
softmax and $\text{scores}\cdot V$entirely within a head, so entirely within a ranknone
$W_O$row parallelone all-reduce

Each rank runs a complete, correct subset of the heads. Softmax normalizes over key positions, not over heads, so no cross-rank reduction is needed inside attention at all.

Two all-reduces per block, forward. One at the end of attention, one at the end of the MLP. With the backward pass adding a mirrored pair, a transformer block costs four all-reduces of the activation tensor per training step.

The bandwidth bill

Each all-reduce moves $2\frac{P-1}{P}$ times the activation tensor $(B,S,d)$ per rank. For a 70B-class model — $d=8192$, $L=80$ — with $S=4096$ and one sequence per rank in BF16, the activation is 64 MiB and the arithmetic is brutal:

$P$Per all-reducePer block (4 collectives)Whole model per stepOn NVLinkOn InfiniBand
264 MiB256 MiB20.0 GiB24 ms429 ms
496 MiB384 MiB30.0 GiB36 ms644 ms
8112 MiB448 MiB35.0 GiB42 ms752 ms
This is why tensor parallelism never crosses a node boundary. The same traffic takes 42 ms on a 900 GB/s intra-node fabric and 752 ms on a 50 GB/s network — an 18-fold penalty, paid on every step. Tensor parallel groups are sized to the number of accelerators sharing a fast interconnect, which in practice means 8.

Note also that this communication is not overlappable in the way DDP's gradient all-reduce is. It sits directly between two dependent matmuls: the next layer cannot start until the sum arrives. Tensor parallelism buys memory with latency you cannot hide.

weight memory / rank
TP traffic / step
ms on NVLink
norm activation / rank
A 70B-class model: $d=8192$, $L=80$, BF16, one sequence per rank. Toggling sequence parallelism replaces each all-reduce with a reduce-scatter and an all-gather, which moves the same bytes but shards the tensors between blocks.

Sequence parallelism

There is a gap in the scheme above. Between the two all-reduces, the tensor is replicated — every rank holds the identical full $(B,S,d)$ activation. The normalization layers, residual adds and dropout all operate on that replicated tensor, so each rank stores an identical copy and performs identical work. With $P=8$ that is eight duplicated copies of the largest activations in the model.

These operations are elementwise across the hidden dimension but independent across sequence positions, so they can be sharded along $S$ instead. The transitions then become:

Norm regionsharded along sequence
→
all-gatherassemble full sequence for the matmuls
→
reduce-scattersum partials and reshard by sequence
Same bytes, less memory. An all-reduce is a reduce-scatter plus an all-gather, so splitting it in two and doing useful work in between costs nothing extra in bandwidth. The activations in the norm regions shrink by $P$. This is one of the rare optimizations that is simply free.

The two edges of the model

Vocabulary-parallel embedding

Shard the $V\times d$ table by vocabulary. Each rank looks up only the ids it owns and returns zeros elsewhere; one all-reduce assembles the true embedding. The table is often the single largest tensor, so this matters.

Vocabulary-parallel cross-entropy

The logits tensor is $(B,S,V)$ — for a 128k vocabulary it can exceed the model itself. Keep it sharded and compute the loss in place: all-reduce the per-shard maximum for numerical stability, then the sum of exponentials. Two scalars per token cross the wire instead of the whole logit tensor.

Rules of thumb

QuestionAnswer
How large should the TP group be?The number of devices on one fast fabric, typically 8. Never larger.
Can $P$ exceed the head count?No. Heads are the unit of splitting; $h$ must be divisible by $P$.
Should I enable sequence parallelism?Yes. Same traffic, strictly less activation memory.
TP or FSDP for memory?FSDP first — its traffic overlaps with compute. Reach for TP when a single layer's activations or weights do not fit.
Where does TP sit in the mesh?Innermost, on the fastest axis. This is the ordering rule the device-mesh note formalizes.
Takeaway

Column parallel shards the output and needs nothing forward; row parallel produces partial sums and needs an all-reduce. Chain them so each MLP and each attention block costs exactly one collective, split that collective in two to get sequence parallelism for free, and keep the whole group on one fast interconnect.

Check yourself