Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Explain ring attention and how it enables very long contexts across multiple devices?
What you need to know
The problem it solves
For very long sequences, one GPU cannot hold the activations. For Llama-3-8B, the KV cache alone at 1M tokens is:
128 KiB per token × 1,048,576 tokens = 128 GiB → more than an 80 GB GPUTraining is worse, because every layer's activations must be kept for the backward pass. You need to split the sequence itself across GPUs — called sequence or context parallelism.
The ring
- Split — the sequence is cut into N contiguous blocks. Device
iholdsQ_i,K_i,V_i. - Compute — each device computes attention of its
Q_iagainst the K/V block it currently has, and updates a running max, sum and output for its rows. - Pass — each device sends its K/V block to the next device and receives one from the previous device. The transfer overlaps with step 2.
- Repeat — after N rounds every Q block has met every K/V block. Divide the output by the running sum.
This numpy simulation with 4 "devices" matches ordinary attention:
1import numpy as np2rng = np.random.default_rng(1)3N, blk, d = 4, 3, 8 # 4 devices, 3 tokens each4Q, K, V = (rng.normal(size=(N * blk, d)) for _ in range(3))5Qs, Ks, Vs = np.split(Q, N), np.split(K, N), np.split(V, N)6m = [np.full(blk, -np.inf) for _ in range(N)]7l = [np.zeros(blk) for _ in range(N)]8acc = [np.zeros((blk, d)) for _ in range(N)]9kv = list(zip(Ks, Vs)) # block each device holds now10for step in range(N):11 for dev in range(N):12 k, v = kv[dev]13 s = Qs[dev] @ k.T / np.sqrt(d)14 m_new = np.maximum(m[dev], s.max(1))15 c, p = np.exp(m[dev] - m_new), np.exp(s - m_new[:, None])16 l[dev] = l[dev] * c + p.sum(1)17 acc[dev] = acc[dev] * c[:, None] + p @ v18 m[dev] = m_new19 kv = kv[-1:] + kv[:-1] # pass K/V to the next device20ring = np.vstack([a / li[:, None] for a, li in zip(acc, l)])21s = Q @ K.T / np.sqrt(d); p = np.exp(s - s.max(1, keepdims=True))22print(np.allclose(ring, (p / p.sum(1, keepdims=True)) @ V)) # TrueWhen the communication is "free"
Each round, a device computes blk × blk scores and sends one K/V block of blk tokens. Compute grows with blk², communication with blk. So with large enough blocks and a fast interconnect (NVLink, InfiniBand), the transfer hides behind the compute. The original paper is Liu, Zaharia and Abbeel (2023).
The causal-mask problem
With a causal mask, the device holding the first block needs only its own K/V, while the last device needs all of them. Plain contiguous blocks leave early devices idle. Striped attention (Brandon et al., 2023) and zigzag layouts give each device a mix of early and late tokens so the work evens out.
Related approaches
- DeepSpeed-Ulysses splits by attention heads using all-to-all communication instead of a ring.
- The Llama 3 paper describes context parallelism for its long-context training stage using an all-gather of K/V rather than a ring.
A real-life example
A code-assistant team continues pretraining an 8B model at 1M tokens so it can read large repositories at once. On one GPU, the K/V alone would need 128 GiB. With ring attention over 8 GPUs, each holds about 131K tokens, or 16 GiB of K/V, plus its share of activations. In each layer, every GPU passes its K/V block for that layer (about 0.5 GiB, one thirty-second of its 16 GiB) around the ring over NVLink while computing on the block it just received. The team uses a zigzag split so that the GPU holding the first part of each file is not idle while others work.
Follow-up questions to expect
- "Does ring attention reduce FLOPs?" — No. It does the same attention math, split across devices; it removes the memory limit, not the quadratic compute.
- "How is it different from tensor parallelism?" — Tensor parallelism splits the weight matrices across GPUs; ring attention splits the sequence. Large runs combine both.
- "Is it used at inference?" — Mostly for training and for prefilling very long prompts; decode has only one new query, so other schemes are used there.