Skip to content

Attention Mechanism

Quick overview Attention is the heart of Transformers and all modern large models. This article clarifies the idea of "soft retrieval" along with QKV scaled dot-product attention, multi-head attention, three types of positional encoding, self-/cross-attention, and causal masking, while also covering engineering essentials like O(n²) complexity, FlashAttention, and KV cache.

Attention Mechanism ​

One-sentence definition: Attention is a mechanism that "aggregates information weighted by relevance" — given a query, it softly selects the most relevant items from a set of candidates to read. It was first used by Bahdanau et al. in 2014 for machine translation (aligning source words), brought into the spotlight by "Attention Is All You Need" (Vaswani et al., 2017), and became the cornerstone of the Transformer architecture and all large models today. See RNNs and Sequence Modeling for the comparison with traditional RNNs.

1. The Attention Idea: Soft Retrieval ​

Think of attention as a soft database query: you have a "question" (query) and a bunch of "entries to look up" (each with a key and value). Traditional hard retrieval is "find the single matching entry and extract its content"; soft retrieval is "score every entry by its match quality, and take a weighted sum of all entries":

Attention = Σᵢ (match(query, keyᵢ) normalized) × valueᵢ

Three intuitions to keep in mind:

  1. Match quality is measured by dot product / similarity, normalized into weights (summing to 1) — i.e., "softly distributing an attention budget."
  2. The output is a weighted sum of values — always differentiable, with gradients flowing smoothly (see Backpropagation and Automatic Differentiation).
  3. It makes no assumption about sequential dependencies: any query can directly "see" any position's key, naturally solving long-range dependencies — exactly the pain point RNNs can't parallelize and struggle to capture (see RNNs and Sequence Modeling).

2. QKV and Scaled Dot-Product Attention ​

Each input token is projected into three groups of vectors via three different linear transformations:

Q = X·W_Q      K = X·W_K      V = X·W_V

The scaled dot-product attention formula:

Attention(Q, K, V) = softmax(Q·Kᵀ / √d_k) · V

Breakdown:

  • Q·Kᵀ: similarity scores between each pair of tokens, shape (n, n), where n is the sequence length.
  • /√d_k scaling: d_k is the key dimension. Why divide? Because the variance of dot products grows linearly with d_k; unscaled softmax would enter saturation regions, with gradients approaching 0. Scaling keeps the logits' variance at O(1) — the same principle of "variance conservation" we repeatedly emphasized in Initialization and Normalization.
  • Softmax normalizes across rows: each query's attention weights sum to 1.
  • Finally, multiply by V to get each position's representation, "fused" by attention weights.

A hand-rolled PyTorch implementation (for teaching, not performance-optimized):

python
import torch, torch.nn.functional as F

def scaled_dot_product_attention(q, k, v, mask=None):
    # q,k,v: (batch, heads, seq, d_k)
    scores = q @ k.transpose(-2, -1) / (k.size(-1) ** 0.5)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    weights = F.softmax(scores, dim=-1)
    return weights @ v

3. The Motivation for Multi-Head Attention ​

A single attention head computes only one kind of "relationship". Multi-head attention (MHA) splits Q/K/V into h pieces each, computes h attention groups in parallel, concatenates them, and projects again:

MultiHead(Q,K,V) = Concat(head₁,…,head_h)·W_O
headᵢ = Attention(Q·W_Qⁱ, K·W_Kⁱ, V·W_Vⁱ)

Three layers of motivation:

  1. Different heads learn different relationship patterns — some attend to syntactic dependencies, some to coreference, some to positional proximity (confirmed by extensive visualization studies, see Interpretability and Fairness).
  2. Multiple low-dimensional subspaces have stronger expressive power than one high-dimensional space — the total computation is roughly the same (each head's dimension is smaller after splitting), but parameters work in parallel across multiple representation subspaces.
  3. Improved optimization stability: multi-head averages out the noise of a single head.

4. Positional Encoding: Absolute, Relative, and RoPE ​

Attention is "permutation-equivariant" — it is insensitive to input order (shuffling the order of Q·Kᵀ gives the same result). To let the model "know the order," we must inject positional information. Three mainstream approaches:

ApproachHow it worksCharacteristicsNotable examples
Absolute positional encodingA vector per position, added to tokensSimple and direct; weak extrapolation to longer sequencesSinusoidal (original Transformer), learnable positional encoding (GPT series)
Relative positional encodingEncodes "relative offset between two positions," injected into attention scoringExplicitly models relative distance, better extrapolationShaw 2018, Transformer-XL
RoPE (Rotary Positional Encoding)Uses rotation matrices to encode position into the angle between Q/K; their dot product naturally contains relative positionCombines the implementation benefits of absolute encoding with the properties of relative encoding, strong length extrapolationLLaMA, Qwen, DeepSeek, and other modern LLMs

Why have modern LLMs almost all shifted to RoPE? The core requirement is length extrapolation: if pretraining uses sequences of length 4k, inference should support 8k/32k. RoPE's positional information enters the dot product as "phase," naturally depending only on relative displacement. Combined with interpolation techniques (NTK-aware, YaRN), it can be extended to longer contexts at low cost. See Large Language Models (LLM) for details.

5. Self-Attention vs. Cross-Attention ​

  • Self-attention: Q, K, V all come from the same sequence — each token interacts with every other token in that sequence. Purpose: capture long-range dependencies within the sequence. Used internally in both Transformer encoder and decoder.
  • Cross-attention: Q comes from one sequence (e.g., the decoder's current state), while K and V come from another sequence (e.g., the encoder's output). Purpose: let the decoder "read" encoder information — machine translation, image-text multimodal alignment (Multimodal Models), all rely on this.

The essential difference: self-attention models "internal relationships," while cross-attention models "alignment relationships between two sequences."

6. Causal Masking ​

Autoregressive generation (predicting one token at a time) requires the model to only see tokens at "the current position and before," otherwise it would cheat (future tokens leaking into the prediction). Implementation is straightforward: fill the upper triangle of the Q·Kᵀ score matrix with -inf; after softmax, those weights naturally become 0:

mask lower triangle = True, upper triangle = False
scores = scores.masked_fill(~mask, float('-inf'))

This is the causal mask — the formal meaning of "attention can only look left" in autoregressive models. It pairs naturally with autoregressive generation and the "next token prediction" objective from Generative Models.

7. O(n²) Complexity and FlashAttention ​

Self-attention's score matrix is (n, n), so both time and memory grow quadratically with sequence length:

Complexity O(n²·d) — at n=8k, that's 6.4×10⁷ scores

This is the Achilles' heel of Transformers for long sequences. Mitigation strategies split into two paths:

  1. Sparse / linear attention: let each query interact only with a subset of keys (local windows, global tokens, sliding windows), e.g., Longformer, BigBird, Linformer; complexity drops to O(n) or O(n log n).
  2. FlashAttention (Dao et al., 2022): doesn't change the attention math, only the IO — computes attention in blocks, avoids writing intermediate results to GPU memory (tiling + online softmax), minimizing data movement between slow GPU memory and fast SRAM. Result: 2–4× speedup, memory reduced from O(n²) to O(n), with numerical equivalence. It has become the default implementation for training all large models on A100/H100; the Triton tutorial even shows it can be reproduced in a few dozen lines.

Rule of thumb for engineering choices: short sequences (<1k) use standard attention; for long sequences, try FlashAttention first; if still insufficient, go for sparse attention.

8. KV Cache: Inference Acceleration ​

During generative inference, only one new token is added per step, yet attention needs to "see" the K and V of all historical tokens. If every step recomputes all of history, complexity grows quadratically with steps — unusably slow.

KV cache: cache the K and V of already-generated tokens in GPU memory, compute Q/K/V only for the new token each step, concatenate with the cache, and then compute attention:

K_cache = concat(K_cache, K_new)     # append each step
scores = Q_new · K_cacheᵀ / √d_k

Effect: generation latency drops from O(t²) to O(t) (t is the number of generated steps). The cost is GPU memory: KV cache size = number of layers × number of heads × sequence length × dimension × 2 × precision × batch. This is why large models have "insufficient GPU memory for max context," and why GQA/MQA (shared K/V across heads) exist — see Large Language Models (LLM).

9. Attention Visualization and Interpretability ​

Attention weights can be visualized directly as heatmaps: which key tokens did a given query token assign high weight to. Classic findings (e.g., coreference resolution, translation alignment) look highly interpretable, but be cautious:

  • High attention ≠ causal importance. Studies (Jain & Wallace 2019) have shown that removing high-attention heads often leaves predictions unchanged — attention weights contain redundancy.
  • Attention is a complex combination across many layers and heads; a single-layer heatmap only reveals "local correlation."

A more reliable approach is to study attention within the framework of mechanistic interpretability (probes, circuits); see Interpretability and Fairness.

10. Trade-offs ​

Trade-offs

Expressive power vs. computational complexity: Full attention has the strongest expressive power but at O(n²); sparse attention has linear complexity at the cost of some long-range modeling (requiring local windows + global tokens to compensate). Prioritize saving compute for long sequences; don't bother for short ones.

KV cache: latency vs. GPU memory: Caching makes generation an order of magnitude faster, but GPU memory grows linearly with context; GQA/MQA trade a small amount of quality for throughput by sharing KV.

RoPE extrapolation vs. training interpolation: RoPE has the best extrapolation but is not infinite; very long contexts also need NTK/YaRN interpolation, which introduces a small precision loss.

Absolute encoding simplicity vs. relative/RoPE generality: Absolute encoding is sufficient for simple scenarios (fixed-length sequences); length extrapolation requires relative encoding. Modern LLMs choosing RoPE isn't cost-free — its implementation and fused kernel complexity are higher.

The attention mechanism makes "relevance" a first-class citizen. It is both a kind of "layer" in Neural Network Fundamentals and the core source of "contextual representations" in Representation Learning and Pretraining. Once you understand attention, the skeleton of Transformers and LLMs becomes clear — next stop: Transformer Architecture.

Further Reading ​

References ​