Fast Attention and Long Context
Ring attention and sequence parallelism
Ring attention splits one long sequence across many devices and passes the keys and values around a circle, so context length grows with the number of machines instead of the memory of one.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Ring attention gives each machine a slice of the text, then passes the slices around a circle.
The dishes at a long table
Twenty people sit at a long table. Twenty dishes are laid out, one in front of each person.
Nobody tries to pile all twenty dishes onto their own plate. Each dish is passed to the left. You take a spoonful from whatever arrives, pass it on, and wait for the next.
After twenty passes, every person has eaten from every dish. No single place at the table ever held more than one dish.
That is ring attention. Each machine holds one slice of the text, and the slices travel around the circle.
The problem it solves
The other methods in this section make attention cheaper on one machine. There is a limit to that.
A million words of text needs more memory for its stored notes than any single graphics card has. No amount of cleverness inside one card fixes a shortage of that size.
So use more cards. Split the text, not the model. Ten cards hold ten times the text; a hundred cards hold a hundred times.
device 1 device 2 device 3 device 4
words 1-1000 1001-2000 2001-3000 3001-4000
| | | |
└──── pass ────>└──── pass ────>└──── pass ────>┘
(and around, back to device 1)
after 4 passes every device has seen all 4000 wordsWhy this is not obvious
The combining step is the hard part. It needs to see all the scores at once to work out the percentages.
The fix is the one from FlashAttention: keep a running total and correct it as new pieces arrive. Each machine keeps its own running total and updates it as each slice comes past.
The answer at the end is exactly the same as if one enormous machine had done it all. Nothing is approximated.
The free lunch that makes it fast
A machine works on the slice in its hands. The next slice is already travelling to it over the network.
If the sum takes longer than the journey, the journey is free. The network cost disappears behind the work. That overlap is why this method is practical rather than a curiosity.
The awkward part nobody mentions first
Text models read left to right. The first word looks at almost nothing; the last word looks at everything before it.
Split the text into four equal blocks. The machine with the last block does far more work than the one with the first. Three machines finish and wait.
The fix is to hand each machine a mix of early and late pieces instead of one continuous block. Boring to implement, and it roughly doubles the useful speed.
Where you have already seen this
- Models advertising a one-million-word context.
- Training runs on hundreds of graphics cards for a single long document.
- Video and genome models, where one example is enormous.
Remember this
- Each device holds a slice of the sequence. Slices are passed around a ring.
- Running totals with corrections make the split answer identical to the whole answer.
- Context grows with the number of devices. Splitting fairly is a real engineering problem.
What to learn next
- Advertised context vs usable context — whether the context you paid to train is the context you get.
- FlashAttention — the tiling argument ring attention lifts onto a network.
- Docker for ML — packaging a distributed training job so it runs the same everywhere.
Developer — Code and libraries.
Setup
pip install numpyWritten against NumPy 1.26.4, Python 3.10. This simulates the ring on one CPU, so you can read the algorithm without a cluster.
The whole algorithm, simulated
import numpy as np
rng = np.random.default_rng(0)
N, d, P = 16, 8, 4 # 16 tokens, head dim 8, a ring of 4 "devices"
S = N // P # 4 tokens per device
Q, K, V = (rng.normal(size=(N, d)).astype(np.float32) for _ in range(3))
scale = 1.0 / np.sqrt(d)
def reference(Q, K, V):
s = (Q @ K.T) * scale
s = s - s.max(-1, keepdims=True)
p = np.exp(s)
return (p / p.sum(-1, keepdims=True)) @ V
# Each device owns one shard of Q, K, V and never sees the whole sequence.
q = [Q[r*S:(r+1)*S] for r in range(P)]
k = [K[r*S:(r+1)*S] for r in range(P)]
v = [V[r*S:(r+1)*S] for r in range(P)]
O = [np.zeros((S, d), np.float32) for _ in range(P)]
m = [np.full(S, -np.inf, np.float32) for _ in range(P)]
l = [np.zeros(S, np.float32) for _ in range(P)]
peak_kv = 0
for hop in range(P): # P steps: each device sees every KV shard once
for r in range(P):
Kb, Vb = k[r], v[r] # after `hop` rotations this is shard (r - hop) % P
peak_kv = max(peak_kv, Kb.nbytes + Vb.nbytes)
s = (q[r] @ Kb.T) * scale
m_new = np.maximum(m[r], s.max(-1))
corr = np.exp(m[r] - m_new)
p = np.exp(s - m_new[:, None])
l[r] = corr * l[r] + p.sum(-1)
O[r] = corr[:, None] * O[r] + p @ Vb
m[r] = m_new
# the ring rotation: every device hands its KV shard to its neighbour
k = k[-1:] + k[:-1]
v = v[-1:] + v[:-1]
ring = np.concatenate([O[r] / l[r][:, None] for r in range(P)])
ref = reference(Q, K, V)
print("ring output matches full attention:", np.allclose(ring, ref, atol=1e-5))
print("max difference:", float(np.abs(ring - ref).max()))
print()
print(f"devices in the ring : {P}")
print(f"tokens held per device : {S} (of {N})")
print(f"largest KV block in flight : {peak_kv:,} bytes")
print(f"whole-sequence KV would be : {K.nbytes + V.nbytes:,} bytes")
print(f"communication steps : {P} rotations, each of one KV shard")
print()
print("what this buys at scale, one head, bf16")
print(f"{'devices':>8}{'seq len':>12}{'per-device KV':>16}{'full KV':>14}")
for P2 in (8, 64, 512):
n = P2 * 32768
per = 2 * (n // P2) * 128 * 2
print(f"{P2:>8}{n:>12,}{per/2**20:>13.1f} MB{2*n*128*2/2**30:>11.1f} GB")ring output matches full attention: True
max difference: 1.7881393432617188e-07
devices in the ring : 4
tokens held per device : 4 (of 16)
largest KV block in flight : 256 bytes
whole-sequence KV would be : 1,024 bytes
communication steps : 4 rotations, each of one KV shard
what this buys at scale, one head, bf16
devices seq len per-device KV full KV
8 262,144 16.0 MB 0.1 GB
64 2,097,152 16.0 MB 1.0 GB
512 16,777,216 16.0 MB 8.0 GBReading the output
The ring result equals full attention to 1.8e-07. This is the property that matters. Ring attention is not an approximation, a windowing scheme or a sparsification. It is the same function, evaluated in a different order across machines.
The running statistics m, l and O are per device. Compare them to the FlashAttention loop: identical structure. Ring attention is FlashAttention's tiling where the outer loop crosses a network instead of a memory hierarchy. That is the whole idea in one sentence.
largest KV block in flight is one shard, not the sequence. Per-device memory is set by shard size. That is why the last table shows a flat 16 MB per device while total sequence length grows to 16 million.
The last column grows and the third does not. Add devices, get context. This is the only technique in this section with that property.
Where the rotation happens for real
In a distributed implementation the list rotation becomes a point-to-point send and receive:
# conceptual, using torch.distributed
next_k = torch.empty_like(k)
req_s = dist.isend(k, dst=(rank + 1) % world_size) # start sending now
req_r = dist.irecv(next_k, src=(rank - 1) % world_size)
out, m, l = attend_block(q, k, v, out, m, l) # compute while it travels
req_s.wait(); req_r.wait()
k = next_kThe isend/irecv before the compute call is the entire performance story. If one block of attention takes longer than one shard's network hop, communication is fully hidden and the ring costs nothing beyond the arithmetic.
What to use instead of writing this yourself
- PyTorch context parallel (
torch.distributed.tensor.experimental.context_parallel) — documented as a prototype, and used bytorchtitan. The PyTorch team reports scaling Llama3-8B to 1M sequence length on 32 H100s, composed with FSDP andtorch.compile. ring-flash-attention— community kernels wrappingflash-attnwith the ring loop.- DeepSpeed-Ulysses — a different sequence-parallel scheme that all-to-alls over the head dimension rather than passing KV around a ring.
The PyTorch implementation offers two variants: an all-gather "pass-KV" algorithm and an all-to-all one. Their benchmarks found all-to-all rarely beats all-gather, and all-gather is the default. Llama 3 training used the all-gather pass-KV form.
The load-balance problem, concretely
With causal masking and contiguous shards, device r computes attention for its queries against shards 0..r. Device 0 does one block of work; device P-1 does P blocks. Average utilisation is about 50%.
The standard fix is zigzag or striped assignment: give each device two half-shards, one from the front of the sequence and one from the back. Every device then holds a mix of early and late queries and the per-hop work evens out.
Any ring implementation that does not mention load balancing is leaving roughly half the throughput on the floor.
Common mistakes
Confusing sequence parallelism with tensor or pipeline parallelism. Tensor parallelism splits the weights of a layer. Pipeline parallelism splits layers across devices. Sequence parallelism splits the tokens. They compose, and they solve different shortages.
Ignoring the causal imbalance. Covered above. Symptom: adding devices scales context but throughput per device drops by half.
Choosing a shard size that leaves the GPU idle. If one block of attention finishes faster than a network hop, communication stops being hidden. Larger shards, fewer devices, or a faster interconnect.
Running a ring across slow links. Inside a node over NVLink this works well. Across nodes over standard Ethernet, the hop dominates and you have built a slow machine. Check your interconnect bandwidth against your per-block compute time before designing around this.
Assuming it helps decoding. Ring attention shines when there are many queries to process at once, which means prefill and training. Single-token decoding has one query row and nothing to overlap.
Try it yourself
Add causal masking to the simulation: at hop h, device r should skip shards from the future entirely. Count the blocks each device computes and confirm the 1, 2, 3, 4 imbalance. Then implement zigzag sharding and watch the counts even out.
What to learn next
- Advertised context vs usable context — whether the context you paid to train is the context you get.
- FlashAttention — the tiling argument ring attention lifts onto a network.
- Docker for ML — packaging a distributed training job so it runs the same everywhere.
Researcher — Mathematics and papers.
Formulation
Liu, Zaharia and Abbeel (2023), Ring Attention with Blockwise Transformers for Near-Infinite Context (arXiv:2310.01889).
Partition the sequence of length $N$ into $P$ contiguous shards of size $S = N/P$, one per device. Device $r$ holds $Q^{(r)}, K^{(r)}, V^{(r)} \in \mathbb{R}^{S \times d}$ and the running triple $(O^{(r)}, m^{(r)}, \ell^{(r)})$.
At hop $j = 0 \dots P-1$, device $r$ holds shard $(r - j) \bmod P$ and applies the online-softmax update from FlashAttention:
$$ m \leftarrow \max(m, \operatorname{rowmax}(S^{(j)})), \quad \ell \leftarrow e^{m_{\text{old}} - m}\ell + \operatorname{rowsum}(e^{S^{(j)} - m}), \quad O \leftarrow e^{m_{\text{old}} - m} O + e^{S^{(j)} - m} V^{(j)} $$
Simultaneously it sends its current $K, V$ to device $(r+1) \bmod P$ and receives from $(r-1) \bmod P$. The paper's claim is exactness plus full overlap: sequences "up to device count times longer" than prior memory-efficient transformers, achieved by "fully overlapping the communication of key-value blocks with the computation of blockwise attention", "without resorting to approximations or incurring additional communication and computation overheads".
The blockwise feed-forward companion is Liu and Abbeel (2023), Blockwise Parallel Transformer for Long Context Large Models (arXiv:2305.19370).
The overlap condition
Per hop, compute is $\Theta(S^2 d)$ FLOPs and communication is $2 S d b$ bytes for bf16 $b = 2$. Communication is hidden when
$$ \frac{2 S^2 d \cdot c}{F} \;\ge\; \frac{2 S d b}{W} \quad\Longleftrightarrow\quad S \;\ge\; \frac{b F}{c W} $$
$F$ is per-device FLOP/s, $W$ is interconnect bandwidth in bytes/s, and $c$ is a constant absorbing achieved efficiency. The shard size must exceed a threshold set by the compute-to-bandwidth ratio of the machine. On an NVLink-connected node this threshold is small; over commodity Ethernet it can exceed any shard you would want, and the ring degenerates into a serial bottleneck.
This is the same inequality as the roofline argument in memory-bound vs compute-bound, promoted one level up the hierarchy: registers, SRAM, HBM, then the network.
Load balancing under causality
With causal masking, device $r$ performs work on $r+1$ of the $P$ shards. Total work is $\sum_{r=0}^{P-1}(r+1) = P(P+1)/2$ block-attentions distributed over $P$ devices with a makespan of $P$. Efficiency is
$$ \frac{P(P+1)/2}{P \cdot P} = \frac{P+1}{2P} \to \frac{1}{2} $$
Striped attention (Brandon et al., 2023) assigns each device a strided subset of token indices rather than a contiguous block, equalising per-hop work. Zigzag assignment pairs shard $r$ with shard $2P-1-r$ to the same device and achieves the same effect with contiguous halves, which is friendlier to positional embeddings. Both roughly double achieved throughput.
Alternatives in the same slot
| Scheme | Communication | Splits over |
|---|---|---|
| Ring attention (pass-KV) | P2P ring, or all-gather in PyTorch's default | sequence |
| DeepSpeed-Ulysses | two all-to-alls per layer | sequence, then heads |
| Tensor parallel | all-reduce per layer | hidden dimension |
Ulysses transposes the sharding so each device owns all tokens for a subset of heads during attention, which makes attention itself local at the cost of two all-to-all collectives per layer. Its parallel degree is capped by the head count; ring attention has no such cap. Production stacks often combine them, using Ulysses within a node and a ring across nodes.
PyTorch's benchmarks report all-gather pass-KV outperforming the all-to-all variant in most configurations, and they default to it. Llama 3's long-context training used the all-gather form.
Practical scope
Ring attention is primarily a training and prefill technique. Decoding presents a single query row per sequence, giving no computation to hide communication behind. Long-context inference instead relies on KV cache sharding with a different communication schedule, or on paged caches spilled to host memory.
The composition rules are the useful takeaway. Ring attention is exact and orthogonal to everything else in this section: it composes with GQA, sliding windows (which reduce the number of hops a device must complete), sparse patterns and cache quantisation. It is the only method here that raises the ceiling rather than lowering the cost, which is why million-token training runs use it alongside, not instead of, the rest.
What to learn next
- Advertised context vs usable context — whether the context you paid to train is the context you get.
- FlashAttention — the tiling argument ring attention lifts onto a network.
- Docker for ML — packaging a distributed training job so it runs the same everywhere.