Context Parallelism & Ring Attention
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$:
The ring
Each device keeps its own queries fixed and cycles the keys and values around a ring:
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$.
| $P$ | Per-device work (blocks) | Max / mean |
|---|---|---|
| 2 | 1, 2 | 1.33× |
| 4 | 1, 2, 3, 4 | 1.60× |
| 8 | 1, …, 8 | 1.78× |
| 16 | 1, …, 16 | 1.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:
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$.
Where it fits among the others
| Context parallel | Tensor parallel | Data parallel | |
|---|---|---|---|
| Splits | sequence positions | inside each matmul | examples |
| Parameters | replicated | sharded | replicated |
| Communication | ring of $K,V$ inside attention only | collectives twice per block | one gradient all-reduce per step |
| Overlappable | yes, fully | no, it is on the critical path | yes |
| Reach for it when | one sequence does not fit | one layer does not fit | you 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.
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
- Derive the online softmax update and show it is exact for two blocks. derivation
- Explain why only attention needs communication and the MLP does not. reasoning
- Show that causal contiguous sharding gives an imbalance of $2P/(P+1)$. derivation
- Prove that zigzag pairing gives every device a load of $2P+1$ units. proof
- Find the condition on $S$, $d$ and bandwidth under which the ring transfer fully hides. analysis