Ring Attention: Distributed Reads with Online Softmax
Partition the sequence across devices
Suppose a sequence of length $n$ is split across $P$ devices. Device $r$ holds roughly $n/P$ query positions and one K/V shard. It computes attention contributions from its local K/V shard, then sends that shard to the next device while receiving another shard from its predecessor.
After the required rounds, each local query has interacted with every allowed K/V shard. The output remains on the device that owns that query. This is sequence/context parallelism; it does not require splitting every model weight matrix across the same devices.
Do not average independently normalized shard outputs
One shard can contain far larger scores than another. Averaging each shard's locally normalized output gives them equal influence regardless of their global softmax mass. Instead, retain each shard's running maximum $m$, denominator $\ell$, and numerator $u$, and use the online-softmax merge.
For an extreme scalar example, shard A has one score 0 and value 0; shard B has one score 10 and value 1. Averaging local outputs gives 0.5. Global softmax gives $e^{10}/(1+e^{10})\approx0.99995$. The merge must account for the relative exponential scales.
A simple forward schedule
- Initialize a per-query running softmax summary on each device.
- Compute contributions from the current K/V shard.
- Merge those statistics into the local summary.
- Transfer K/V to the next device and receive the next shard, overlapping where valid.
- After all allowed shards have contributed, normalize $u/\ell$ for the local queries.
The ring is one communication topology for presenting all shards to all query owners. Double buffering can let one buffer participate in computation while another receives data. Synchronization must ensure a transfer finishes before a consumer reads its buffer.
Causal masks depend on global positions
A query on one device may be forbidden from reading keys on another device because they lie in the future. Local row/column indices are insufficient; shard metadata must identify their global positions.
Contiguous causal sharding can create imbalance: early queries have few allowed keys while late queries have many. Alternative assignments, including paired early/late chunks, can distribute work more evenly. They require correspondingly careful position and ownership bookkeeping. The existing context-parallelism chapter develops these layouts further.
Separate local compute from network traffic
Ignoring masks, each device processes $n/P$ queries against $n$ keys, so attention arithmetic per device is of order $n^2d/P$. Each device sends and receives K/V shards over approximately $P-1$ transfers, with shard size proportional to $nd/P$. Total communicated volume per device is therefore of order $nd$ for large $P$, before protocol details.
As $P$ grows, compute per device shrinks, but network latency and bandwidth do not disappear. If a shard's matrix work is shorter than its transfer time, communication cannot be fully hidden. A good overlap diagram does not prove a good speedup for every context length and interconnect.
Backward must aggregate shared dependencies
A local query contributes gradients to keys and values on multiple remote shards. Those contributions must be accumulated and returned to the owners of the corresponding K/V data. Recomputing probability tiles can save memory, but it still requires the same score scaling, masks, and any dropout choices.
The distributed operation should be checked against a single-device reference on a small sequence, including gradients. Matching only forward output misses errors in ownership and reductions that can silently corrupt training.
What is preserved, what is distributed?
In its dense form, Ring Attention preserves the softmax function up to numerical differences while distributing memory and computation. It does not make dense attention's global arithmetic linear in $n$. It can compose with tiled kernels and sparse patterns; each additional technique changes a different part of the cost.
The primary reference is Ring Attention with Blockwise Transformers. Its title's long-context motivation should be read alongside the actual hardware and workload constraints.
Try it: Why is a running numerator vector alone insufficient to merge two attention shards?
The two numerators may be expressed relative to different score maxima, and their denominator masses differ. The maximum and denominator are needed to rescale and normalize their contributions consistently.
Reduce sequential decoding calls
Speculative decoding tackles a separate bottleneck: how often the expensive target model must be invoked while generating tokens.