← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 12 · Chapter 23

Context Parallelism & Ring Attention

Data, tensor and pipeline parallelism all leave one sequence intact. When a single sequence is long enough that its activations alone exhaust a device, the only axis left is the sequence itself. Splitting it works because softmax, despite appearing to need all its inputs at once, can be computed incrementally without error.

The last remaining axis

From the FLOPs note, attention's share of compute is $S/(6d+S)$, and its activation footprint grows with $S$ as well. Past a few hundred thousand tokens, one sequence in one layer no longer fits on one device — and no amount of data parallelism helps, because the indivisible unit is a single example.

Context parallelism shards along sequence position: device $p$ owns tokens $[pS/P,\ (p{+}1)S/P)$. Every device holds the whole model, so parameters are replicated exactly as in data parallelism.

What is easy

Everything pointwise. Norms, MLPs, residual adds and the projections $W_Q,W_K,W_V$ act on each position independently, so each device just processes its own tokens. No communication at all.

What is hard

Attention alone. Query $i$ must attend to every key $j\le i$, and most of those live on other devices.

Softmax is decomposable

The obstacle looks fundamental: softmax normalizes over all keys, so you seemingly cannot start until you have seen them all. But the normalizer is just a sum, and sums can be accumulated — you only need to handle the numerical stabilization carefully.

Process keys in blocks, maintaining three running quantities: the maximum logit $m$, the sum of exponentials $\ell$, and the unnormalized output $O$. On seeing a block with local max $m_b$ and local sum $\ell_b$:

$$m'=\max(m,m_b),$$
$$\ell'=\ell\,e^{\,m-m'}+\ell_b\,e^{\,m_b-m'},$$
$$O'=\frac{\ell\,e^{\,m-m'}\,O+e^{\,m_b-m'}\,P_bV_b}{\ell'}.$$
This is exact, not an approximation. The correction factors $e^{m-m'}$ rescale everything accumulated so far onto the new maximum. Processing all keys in one block or in a thousand blocks gives bit-comparable results. I verified this numerically while writing this note: block-wise and direct attention agree to $2\times10^{-16}$, which is floating-point noise. This same recurrence is what makes FlashAttention possible; ring attention simply distributes the blocks across devices instead of across shared-memory tiles.

The ring

Each device keeps its own queries fixed and cycles the keys and values around a ring:

Ring attention on device $p$
1hold $Q_p$; initialize the local $K,V$ block to $K_p,V_p$
2initialize $m=-\infty$, $\ell=0$, $O=0$
3for step $=0,\dots,P-1$:
4start sending the current $K,V$ block to device $p+1$, and receiving from $p-1$
5meanwhile, attend $Q_p$ against the block in hand and update $m,\ell,O$
6wait for the transfer, then swap in the received block
7after $P$ steps every query has seen every key; $O$ is the exact attention output
The communication hides completely, in principle. Step 4 is issued before step 5 and waited on after, so the transfer of one $K,V$ block overlaps the attention computation on the previous one. Attention on a block costs $O(S^2/P^2 \cdot d)$ while the transfer costs $O(S/P \cdot d)$ bytes, so as $S$ grows compute outpaces communication and the ring becomes free. Long context is exactly the regime where this holds.

The causal load-imbalance problem

Everything above assumed full attention. Add a causal mask and the arithmetic becomes badly unfair. With contiguous shards, device $p$ owns the tokens at positions $[pS/P,(p{+}1)S/P)$, and those queries may only attend to shards $0$ through $p$. So device $0$ does one block of work, device $1$ does two, and device $P-1$ does $P$.

$$\text{work}_p\propto p+1,\qquad \frac{\max_p\text{work}_p}{\text{mean}}=\frac{2P}{P+1}.$$
$P$Per-device work (blocks)Max / mean
21, 21.33×
41, 2, 3, 41.60×
81, …, 81.78×
161, …, 161.88×

Since the ring is synchronous, everyone waits for the busiest device, so nearly half the fleet's attention capacity is wasted as $P$ grows.

The zigzag fix

Nothing requires a device's tokens to be contiguous. Split the sequence into $2P$ chunks instead of $P$, and give device $p$ one chunk from the front and one from the back:

$$\text{device }p\ \text{owns chunks}\ \{p,\ 2P-1-p\}.$$

An early chunk is cheap and a late chunk is expensive, and pairing them makes every device's total identical. For $P=4$ the pairs are $(0,7),(1,6),(2,5),(3,4)$, and each device's load is exactly $9$ units — perfectly balanced, which I confirmed numerically for several $P$.

Why the pairing works. Chunk $c$ costs about $c+1$ units, so device $p$'s total is $(p+1)+(2P-p)=2P+1$, independent of $p$. The imbalance vanishes exactly, at zero communication cost — only the token-to-device mapping changed.
busiest device
mean work
imbalance
capacity wasted
Work is measured in causal key-blocks each device must process. The ring is synchronous, so throughput is set by the busiest device and everything above the dashed mean line is idle time elsewhere.

Where it fits among the others

Context parallelTensor parallelData parallel
Splitssequence positionsinside each matmulexamples
Parametersreplicatedshardedreplicated
Communicationring of $K,V$ inside attention onlycollectives twice per blockone gradient all-reduce per step
Overlappableyes, fullyno, it is on the critical pathyes
Reach for it whenone sequence does not fitone layer does not fityou have spare devices

Context parallelism composes cleanly with the others because it touches a different axis: shard the sequence across one mesh dimension, the heads across another, the layers across a third, and the batch across a fourth. The device-mesh note puts all four together.

It does not reduce parameter memory at all. Every device still holds the full model. Context parallelism is purely an activation-memory technique, so it is always combined with something else — usually FSDP or tensor parallelism — rather than used alone.
Takeaway

Online softmax makes attention exactly decomposable over key blocks; a ring circulates those blocks so the transfer hides behind compute; and pairing an early chunk with a late chunk removes the causal imbalance exactly. That is the whole of context parallelism.

Check yourself