๐Ÿ’ฌ NLP & Transformers ยท Lecture 29 of 29

Efficient Transformers: Sparse Attention, Linear Attention and FlashAttention

Self-attention's quadratic cost limits context length. We survey sparse and local attention, low-rank and kernel-based linear attention, IO-aware FlashAttention, and alternatives such as state-space models โ€” with their trade-offs.

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,
$$ \text{Attn}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) \approx \frac{\phi(\mathbf{Q})\big(\phi(\mathbf{K})^\top\mathbf{V}\big)}{\phi(\mathbf{Q})\big(\phi(\mathbf{K})^\top\mathbf{1}\big)} $$

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.

python
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.
JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ’ฌ NLP & Transformers

Text-to-Speech: From Concatenation to Neural Voices

Text-to-speech systems turn written text into natural-sounding speech. We cover the TTS pipeline, text normalisation and phonemes, acoustic models like Tacotron and FastSpeech, neural vocoders, end-to-end and zero-shot voice models, evaluation, and voice-cloning ethics.

Intermediateโฑ 5 min#189
๐Ÿ’ฌ NLP & Transformers

Automatic Speech Recognition: From HMMs to Whisper

Speech recognition converts audio into text. We cover audio features and spectrograms, the classical HMMโ€“GMM pipeline, end-to-end neural models with CTC and attention, self-supervised wav2vec 2.0, Whisper, and evaluation with word error rate.

Intermediateโฑ 5 min#188
๐Ÿ’ฌ NLP & Transformers

Dialogue Systems and Chatbots: From Rules to LLM Assistants

We compare rule-based, task-oriented and open-domain dialogue systems; cover intent detection, slot filling and dialogue state tracking; and show how LLM-based assistants with retrieval and tools are designed, evaluated and deployed safely.

Intermediateโฑ 5 min#187