A Visual Guide to Parallel Transformers
Illustrating how a transformer’s forward and backward passes are sharded across devices — data, tensor, context, pipeline, and expert parallelism, in the notation of the JAX scaling book.
How to read the notation #
Everything here is written in the scaling book’s sharding notation. The devices form a mesh — Mesh({'X': 4}) is four devices along one axis named X — and In[BX, T, D] means the batch dimension B is sharded four ways over X, so each device holds a [B/4, T, D] slice. A tensor with no subscripts anywhere is fully replicated. A trailing {UX} marks an unreduced partial sum — each device holds a full-shaped tensor that is only its contribution to the true result, and the shards must still be summed over X.
How to read the figures #
Each figure is a small machine you can drive. Model depth runs left to right: the input enters at In, passes through each layer’s attention and MLP stations, and leaves at Out. Devices are horizontal lanes, stacked vertically — when devices communicate, tensors fly vertically between lanes. Weights are blue, activations amber, keys and values teal, and gradients rose. Hover any tensor for its shape, sharding, and memory at our nominal sizes (B=8, T=128, D=1024, F=4096, bf16).
Data parallelism #
The simplest strategy: replicate the weights everywhere and shard the batch — each device gets In[BX, T, D], a quarter of the sequences, and runs the whole model on them. The forward pass needs zero communication; step through it and watch nothing ever cross between lanes. The price appears in the backward pass: every device computes a weight gradient from only its shard of the batch, so the gradients are unreduced partial sums — dW[D, F]{UX} — and each one must be AllReduced across the data axis before the optimizer can step. That per-weight AllReduce is the classic data-parallel gradient sync.
Fully-sharded data parallelism (ZeRO-3) #
Data parallelism replicates every weight four times — FSDP refuses to pay that memory. Weights are sharded along the same data axis (Win[DX, F], Wout[F, DX]), and each one is AllGathered just in time, used for its matmul, and immediately discarded — watch the station header switch to the gathered shape and back. The backward pass pays the gather again (the weight is long gone), and the gradient flows the other way: each dW is a partial sum that ReduceScatters back to exactly the shard layout the weights live in — memory of a shard, communication of a gather, in both passes.
Tensor parallelism #
Tensor parallelism shards the feature dimensions: attention heads across devices (Wqkv[D, HY], Wo[HY, D]) and the MLP’s hidden width (Win[D, FY], Wout[FY, D]). Activations travel sharded on D. Every block then AllGathers on the way in, ReduceScatters on the way out. The first matmul needs the full D, so the sharded activations are gathered; the second matmul contracts a sharded dimension, so each device is left holding an unreduced partial sum — the dashed {UY} tensor — which the ReduceScatter resolves while re-sharding for the next block.
Context parallelism #
Long sequences don’t fit on one device, so shard the sequence: In[B, TX, D], each device holding a quarter of the tokens. The MLP never notices — tokens are independent there. Attention is where every query must see every key and value, so K and V are AllGathered across the sequence shards while Q stays local. The gathered K/V are used and dropped and the backward pass re-gathers them — the same trade FSDP makes with weights. During the backward pass every device produces gradient contributions for all tokens’ keys and values, so dK and dV are partial sums that ReduceScatter back over the sequence — and the weight gradients AllReduce over the context axis, exactly like data parallelism sums over the batch.
Expert parallelism #
In a mixture-of-experts layer the MLP becomes four experts, one per device (Win[EZ, D, F]), and each token is routed to one of them. The figure’s tokens are small squares colored by their assigned expert: the AllToAll dispatch is the moment every device sends every other device the tokens that belong to it — watch the colors sort themselves into lanes. The experts run an ordinary MLP on their guests, and a second AllToAll sends every token home. The backward pass mirrors it: token gradients AllToAll out to the experts, expert weight gradients stay local (each expert owns its weights outright), and the attention weight gradients AllReduce over the batch axis as in data parallelism.
Pipeline parallelism #
Pipeline parallelism shards the layers: stage 0 owns the first layer, stage 1 owns the second, and activations hop across the boundary with a single point-to-point send — the cheapest communication in this whole article. The catch is idleness: with one batch, each stage would wait on the other, so the batch splits into four microbatches that chase each other through the pipe. The staircase is the pipeline filling, the empty corners are the bubble. Switch to + Backward and the gradients flow right-to-left through the same pipe — each backward cell is drawn twice as wide because backward costs roughly twice the compute, which is why the bubble grows during training.