Self-attention compares every token with every other token: its time and memory grow quadratically with sequence length $T$. Doubling the context from 8K to 16K tokens quadruples attention cost. Yet many applications need long contexts โ entire books, legal case files, codebases, long conversations, genomic sequences. A large research effort has produced ways to make transformers efficient. Some change the mathematics; the most influential simply changed how exact attention is computed.
Where the cost comes from#
For each head, computing $\mathbf{Q}\mathbf{K}^\top$ costs $O(T^2d)$ time and materialising the $T \times T$ score matrix costs $O(T^2)$ memory. For $T = 32{,}768$, a single fp16 attention matrix per head is about 2 GB.
Family 1: sparse and local attention#
Restrict each token to attend to a subset of positions:
- Sliding-window (local) attention: each token attends to $w$ neighbours โ cost $O(Tw)$. Stacking layers expands the receptive field, like convolutions. Used (sometimes alternating with full attention) in several modern LLMs.
- Dilated windows: skip positions to cover longer ranges.
- Global tokens: a few tokens (e.g. [CLS] or task tokens) attend to and are attended by everything.
- Longformer (2020) combined sliding windows, dilation and global attention; BigBird (2020) added random attention and proved such sparse patterns retain theoretical expressivity.
- Block-sparse patterns (Sparse Transformer, 2019) for images, audio and long text.
Trade-off: efficient, but fixed sparsity can miss important long-range interactions.
Family 2: low-rank and kernel (linear) attention#
- Linformer (2020): project keys and values along the sequence dimension to a fixed size $k$, giving $O(Tk)$ โ assumes the attention matrix is approximately low rank.
- Kernel-based linear attention: replace the softmax kernel $\exp(\mathbf{q}^\top\mathbf{k})$ with a feature map $\phi(\mathbf{q})^\top\phi(\mathbf{k})$. Then, by associativity,
which costs $O(Td^2)$ โ linear in $T$. Performer (2020) approximated softmax attention with random features. In causal form, linear attention becomes a recurrence with a fixed-size state, enabling constant-memory generation.
Trade-off: approximations have generally underperformed exact softmax attention on language modelling quality at equal scale, especially for tasks requiring precise retrieval from context.
Family 3: exact attention, computed smarter โ FlashAttention#
Dao et al. (2022) observed that attention on GPUs is bottlenecked not by arithmetic but by memory traffic between slow high-bandwidth memory (HBM) and fast on-chip SRAM. FlashAttention:
- splits $\mathbf{Q}$, $\mathbf{K}$, $\mathbf{V}$ into tiles that fit in SRAM;
- computes attention tile by tile using an online softmax (maintaining running maxima and normalisers, so the softmax can be computed incrementally without seeing the whole row);
- never writes the full $T \times T$ matrix to HBM, and recomputes what is needed in the backward pass.
The result is exact attention with memory linear in $T$ and substantial wall-clock speed-ups. FlashAttention-2 and -3 improved parallelism and hardware utilisation further. It is now standard; in PyTorch, F.scaled_dot_product_attention dispatches to fused kernels automatically.
import torch, torch.nn.functional as F, time
def naive_attention(q, k, v):
s = q @ k.transpose(-2, -1) / q.size(-1) ** 0.5
mask = torch.triu(torch.ones(s.shape[-2:], dtype=torch.bool, device=q.device), 1)
return s.masked_fill(mask, float("-inf")).softmax(-1) @ v
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
for T in [1024, 4096]:
q = k = v = torch.randn(1, 8, T, 64, device=device, dtype=dtype)
for name, fn in [("naive", naive_attention),
("fused", lambda q, k, v: F.scaled_dot_product_attention(q, k, v, is_causal=True))]:
t0 = time.time(); fn(q, k, v)
if device == "cuda":
torch.cuda.synchronize()
print(f"T={T:>5} {name:<6} {1000 * (time.time() - t0):7.1f} ms")The online softmax trick#
To compute $\text{softmax}(\mathbf{x})$ over a row seen in chunks, keep a running maximum $m$ and running sum $\ell$. When a new chunk with maximum $m'$ arrives, rescale: $\ell \leftarrow \ell e^{m - m_{\text{new}}} + \sum e^{x_i - m_{\text{new}}}$. The same rescaling applies to the running weighted sum of values. This algebraic trick, combined with tiling, is the heart of FlashAttention.
Family 4: alternatives to attention#
- State-space models (SSMs): S4 (2021) and Mamba (2023) model sequences with structured linear recurrences that can be computed in parallel (as convolutions or scans) during training and as constant-memory recurrences during generation. Mamba adds input-dependent (selective) parameters, closing much of the quality gap with transformers on language.
- Linear RNNs (RWKV, RetNet, xLSTM variants, gated linear attention) pursue similar goals.
- Hybrids interleave a few attention layers with many SSM or linear layers, aiming for transformer-level recall with lower cost.
Other practical techniques for long context#
- KV-cache optimisations: grouped-query attention, cache quantisation, paged memory (vLLM's PagedAttention), evicting less important tokens.
- Ring/sequence parallelism: distribute very long sequences across GPUs.
- Position-encoding scaling (RoPE interpolation, YaRN) to extend context.
- Retrieval instead of length: often it is cheaper and more reliable to retrieve relevant passages than to put everything in the context.