Why one device runs out of room
Attention over a length-N sequence forms scores S = QK^T / √d_k of shape [N, N]. FlashAttention already taught us not to store that matrix: it tiles the computation and streams blocks through fast on-chip memory, so attention costs O(N) memory instead of O(N^2). But tiling does nothing about the other linear-in-N cost — the activations Q, K, V and the residual stream are each [N, d], and for a million-token context those tensors alone dwarf a single device’s memory before you compute a thing.
So the constraint that forces a ring is not the N^2 matrix (FlashAttention handles that) but the N·d activations. A sequence that cannot live on one device must be split across devices — and then every query still needs keys living on every other device. Ring Attention is the communication pattern that meets that need without ever gathering the full sequence anywhere.
Sharding the sequence: what each device holds
Cut the sequence into P contiguous blocks of size b = N/P, one per device. Device i owns block i: its Q_i, K_i, V_i, each of shape [b, d] — sequence parallelism, a split along the token axis.
The asymmetry that makes the ring work: queries stay home. Device i only ever computes outputs for its own query rows, so Q_i never moves. The keys and values must travel, because each query needs q·k for every key. Rather than broadcast all K/V at once — reassembling the full sequence in memory and defeating the purpose — the ring moves one block at a time, so a device ever holds only its own block plus one visiting block: two, not P.
The online-softmax identity that makes streaming legal
Attention’s output for one query is Σ_j softmax(s_j)·v_j = (Σ_j exp(s_j - m)·v_j) / (Σ_j exp(s_j - m)), where m is any constant — taken as max_j s_j for numerical safety, as the subtraction cancels in the ratio. That invariance is why attention can be computed incrementally: if a later block reveals a bigger maximum, you can retroactively fix up everything computed so far.
Keep three running quantities per query — the max m, the denominator l = Σ exp(s - m), and the unnormalized output o = Σ exp(s - m)·v. When a block arrives, raise m if its local max is bigger and rescale the old l and o by α = exp(m_old - m_new) before adding the block — correcting the earlier terms as if they had used the new max all along. This is exactly the FlashAttention accumulator; the ring just feeds it blocks arriving over the network instead of from HBM.
Per query q, keep running state (m, l, o):
m = running max of scores
l = running Σ exp(s - m) # denominator
o = running Σ exp(s - m)·v # vector, length d
For each incoming K/V block:
s_j = q · k_j / √d_k # this block’s scores
m' = max(m, max_j s_j) # new running max
α = exp(m - m') # rescale old state
l' = α·l + Σ_j exp(s_j - m')
o' = α·o + Σ_j exp(s_j - m')·v_j
(m, l, o) = (m', l', o')
After the final block: out = o / l
The ring schedule: passing K/V around
Arrange the devices in a logical ring: 0 → 1 → … → P-1 → 0. The computation runs P steps. On step t, device i holds the K/V block from device (i - t) mod P; it folds that block’s partial attention into its accumulator and — at the same time — passes the block to (i + 1) mod P while receiving the next from (i - 1) mod P.
After P steps each device has seen all P K/V blocks once, so Q_i has attended to the entire sequence. Every device sends and receives exactly one block per step, so traffic is uniform — no hotspots, no all-gather storm — which lets the ring scale to hundreds of devices.
Accumulation math, step by step
The schedule and the accumulator combine into the per-device inner loop, run simultaneously on every device’s own Q_i.
state = (m = -∞, l = 0, o = 0) # per query row in Q_i
kv = (K_i, V_i) # start with your own block
for t in 0 .. P-1:
if t < P-1: async send(kv → i+1), recv(next ← i-1)
S = Q_i · kv.K^T / √d_k # [b, b] scores
state = online_softmax(state, S, kv.V) # rescale + add
if t < P-1: wait(send, recv); kv = next
O_i = state.o / state.l # [b, d] final outputs
The send/recv is launched before the matmul and waited on after, so the next block’s transfer overlaps the current block’s compute. The online_softmax call is the rescaling recurrence above. Nothing here materializes an N×N matrix — the largest tensor any device touches is [b, b].
A worked numeric example
To see the rescale fire, follow one query whose four keys arrive as two blocks, arranged so the second block holds the larger score and forces a real correction. Values are scalars (d = 1) to keep the arithmetic legible; the vector case is identical component-wise.
One query, 4 keys arriving as 2 blocks (/√d_k folded into s):
Block A: s = [1, 2], v = [1, 2]
Block B: s = [3, 0], v = [3, 4]
Step A (m = -∞, l = 0, o = 0)
m' = 2, α = 0
exp(s - 2) = [0.368, 1.000]
l = 1.368 o = 0.368·1 + 1.000·2 = 2.368
Step B (bigger max → real rescale)
m' = 3, α = exp(2 - 3) = 0.368
exp(s - 3) = [1.000, 0.050]
l = 0.368·1.368 + (1.000 + 0.050) = 1.553
o = 0.368·2.368 + (1.000·3 + 0.050·4) = 4.070
out = o / l = 4.070 / 1.553 = 2.621
Direct softmax over [1, 2, 3, 0] = 2.621 ✓
Step A commits a provisional answer believing the max is 2; step B finds a 3, lifts the max, and multiplies the stored l and o by α = 0.368 before adding its terms. That one scalar multiply is the entire cost of processing blocks out of order, yet the result lands exactly on the dense-softmax answer — order-independence is what makes the ring exact, not an approximation.
Memory: O(1) extra per device
Tally what a device stores: its own shard Q_i, K_i, V_i (3·b·d), at most two K/V blocks at once (~2·b·d), plus an accumulator of b·d. Every term is proportional to the block size b = N/P, and none grows with the total context N once you hold b fixed and add devices.
That is the headline result: the extra memory to reach across the whole sequence, beyond a device’s own shard, is one visiting block — O(1) in block count, O(b) in elements, and constant in N. Double the context and the device count together and per-device memory is unchanged, so the ceiling is aggregate pod memory, not any single chip — which is why the ring enables near-arbitrary context.