← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 15 · Chapters 18, 25

Device Meshes & Combining Parallelism

Every strategy so far splits a different axis, so they compose. A device mesh is the bookkeeping that makes composing them tractable: view the flat list of ranks as an $n$-dimensional grid, give each parallelism strategy one dimension, and read off its process group as a line through the grid.

Ranks as coordinates

A job launches $W$ processes numbered $0$ to $W-1$. That flat numbering is useless for reasoning. Instead, factor it:

$$W=D\times P\times C\times T,$$

for data, pipeline, context and tensor degrees, and give every rank a coordinate $(d,p,c,t)$. With the tensor axis varying fastest:

$$\boxed{\text{rank}=\big((d\cdot P+p)\cdot C+c\big)\cdot T+t.}$$
A process group is a line along one axis. The tensor-parallel group of a rank is every rank sharing its $(d,p,c)$ and differing in $t$. The data-parallel group shares $(p,c,t)$ and differs in $d$. Each rank belongs to exactly one group per axis, and the collectives of different strategies act on disjoint communicators.

Concretely, for 512 devices arranged as $D{=}16$, $P{=}4$, $C{=}1$, $T{=}8$:

GroupRanks (for rank 0)StrideSpans
Tensor parallel0–71, contiguousone node
Pipeline parallel0, 8, 16, 2484 nodes
Data parallel0, 32, 64, 96, …32the whole cluster

That the tensor-parallel group comes out as ranks 0–7 — contiguous, and therefore co-located on one machine — is not luck. It is the direct consequence of putting $t$ last in the mapping, and it is the single most important decision in the whole layout.

The ordering rule

Cluster interconnects are hierarchical: devices inside a node share a fabric an order of magnitude faster than the network between nodes. The mapping above makes the innermost axis contiguous, so:

Order the axes by communication intensity, chattiest innermost. The innermost axis gets the fastest links; the outermost tolerates the slowest.
PositionAxisTraffic per stepOverlappable?
innermostTensorvery high — two collectives per block, on the critical pathno
Contexthigh, but only inside attentionyes, fully
Pipelinelow — point-to-point activations at stage boundariespartly
outermostDataone gradient all-reduce per stepyes

Tensor parallelism goes innermost because its collectives sit directly between dependent matmuls and cannot be hidden. Data parallelism goes outermost because its single all-reduce overlaps with the backward pass. Getting this backwards — data parallel inside the node, tensor parallel across nodes — is one of the most expensive misconfigurations available, and it produces a run that works correctly and is several times slower than it should be.

world size
nodes of 8
model state divisor
TP crosses nodes?
Each outlined block is one 8-device node. Highlighting shows which ranks share a process group with rank 0 along the selected axis.

What each axis divides

The axes divide different resources, which is exactly why combining them works:

AxisParametersOptimizer stateActivationsBatch
Data (DDP)———÷ $D$
Data (FSDP)÷ $D$÷ $D$—÷ $D$
Pipeline÷ $P$÷ $P$÷ $P$ roughly—
Tensor÷ $T$÷ $T$÷ $T$ with sequence parallelism—
Context——÷ $C$—
Context parallelism divides no parameters at all, and plain DDP divides nothing but the batch. Neither is a solution to a model that does not fit; each must be paired with one of the others.

A sizing recipe

1

Set $T$ to fit a layer, and no larger. Bounded by the devices on one fast fabric, so at most 8. Use the smallest value that makes a single block's weights and activations fit.

2

Set $C$ from the sequence length. Only if one sequence's activations still do not fit after tensor parallelism. Zero-cost when unused, so leave it at 1 by default.

3

Set $P$ to fit the model across nodes. Pipelining is the cheapest way to cross a slow network. Keep $P$ small enough that $m\ge 9(P-1)$ microbatches remain affordable.

4

Give everything left to $D$. Data parallelism converts spare devices into throughput and is the axis that scales furthest.

For 512 devices and a 70B model: $T=8$ fills a node, $C=1$, $P=4$ splits the model across four nodes, leaving $D=16$. Model state is divided by $P\cdot T=32$ before FSDP does anything further along the data axis.

Pitfalls

MistakeSymptom
TP group spanning two nodesCorrect results, throughput several times lower than expected. Check whether $T$ divides the devices per node.
Global batch not divisible by $D\times m$Ranks disagree on step count and the job hangs at a collective.
Loss averaged over ranks with unequal token countsSilently wrong gradients — the trap from the autograd note.
Gradient clipping computed per rankDifferent clip coefficients on different ranks, so replicas drift apart.
Checkpoint saved with one mesh, resumed with anotherShard shapes mismatch. Save in a layout-independent form.
Sanity check any layout in one line: $D\times P\times C\times T$ must equal the world size, and $T$ must divide the devices per node. Most layout bugs violate one of those two.
Takeaway

Factor the world size into one dimension per strategy, order them by communication intensity with tensor parallelism innermost and data parallelism outermost, and each strategy's process group becomes a line through the grid. The ordering is not a convention — it is what puts the unhideable traffic on the fastest links.

Check yourself