Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Explain ring attention and how it enables very long contexts across multiple devices?


K/V blocks travelling around four GPUsGPU 0 holdsQ0, K/V block 0GPU 1 holdsQ1, K/V block 1GPU 2 holdsQ2, K/V block 2GPU 3 holdsQ3, K/V block 3compute, thenpass K/V onsends to GPU 0After 4 rounds every query block has met every K/V block; running max and sum merge the pieces exactly.
Queries stay put and keys and values move, so a 1M-token context splits into slices each GPU can hold.

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:

Text
128 KiB per token × 1,048,576 tokens = 128 GiB  → more than an 80 GB GPU

Training 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

  1. Split — the sequence is cut into N contiguous blocks. Device i holds Q_i, K_i, V_i.
  2. Compute — each device computes attention of its Q_i against the K/V block it currently has, and updates a running max, sum and output for its rows.
  3. 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.
  4. 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:

Python
import numpy as nprng = np.random.default_rng(1)N, blk, d = 4, 3, 8                          # 4 devices, 3 tokens eachQ, K, V = (rng.normal(size=(N * blk, d)) for _ in range(3))Qs, Ks, Vs = np.split(Q, N), np.split(K, N), np.split(V, N)m = [np.full(blk, -np.inf) for _ in range(N)]l = [np.zeros(blk) for _ in range(N)]acc = [np.zeros((blk, d)) for _ in range(N)]kv = list(zip(Ks, Vs))                        # block each device holds nowfor step in range(N):    for dev in range(N):        k, v = kv[dev]        s = Qs[dev] @ k.T / np.sqrt(d)        m_new = np.maximum(m[dev], s.max(1))        c, p = np.exp(m[dev] - m_new), np.exp(s - m_new[:, None])        l[dev] = l[dev] * c + p.sum(1)        acc[dev] = acc[dev] * c[:, None] + p @ v        m[dev] = m_new    kv = kv[-1:] + kv[:-1]                    # pass K/V to the next devicering = np.vstack([a / li[:, None] for a, li in zip(acc, l)])s = Q @ K.T / np.sqrt(d); p = np.exp(s - s.max(1, keepdims=True))print(np.allclose(ring, (p / p.sum(1, keepdims=True)) @ V))   # True

When 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.