💬 NLP & Transformers · Lecture 14 of 29

Self-Attention in Depth: Intuition, Complexity and Variants

A deeper look at self-attention — what attention heads learn, the geometry of queries and keys, computational complexity, causal masking, KV caching, and efficient variants such as multi-query and grouped-query attention.

In the previous lecture we assembled the Transformer. Now we zoom into its core operation. Self-attention is deceptively simple — three matrix multiplications and a softmax — but understanding its behaviour, costs and modern variants is essential for anyone building or deploying transformer models.

Self-attention as learned, dynamic routing#

For an input sequence $\mathbf{X} \in \mathbb{R}^{T \times d}$:

$$ \mathbf{Q} = \mathbf{X}\mathbf{W}^Q, \quad \mathbf{K} = \mathbf{X}\mathbf{W}^K, \quad \mathbf{V} = \mathbf{X}\mathbf{W}^V, \qquad \mathbf{A} = \text{softmax}\left(\frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d_k}}\right), \quad \mathbf{Y} = \mathbf{A}\mathbf{V} $$

Contrast with other layers:

  • A fully connected layer mixes features with fixed weights.
  • A convolution mixes neighbouring positions with fixed weights.
  • Self-attention mixes all positions with weights $\mathbf{A}$ that are computed from the input itself — the connectivity pattern changes for every sentence.

Row $t$ of $\mathbf{A}$ is a probability distribution saying how much token $t$ reads from each other token.

Why separate queries, keys and values?#

Using the raw embeddings for everything would make attention symmetric ($\mathbf{x}_i^\top\mathbf{x}_j = \mathbf{x}_j^\top\mathbf{x}_i$) and force each token to attend most to itself. Separate projections let a token ask for one kind of information (query) while advertising another (key) and providing a third (value). A verb's query might seek its subject; a noun's key might advertise "I am a plausible subject".

Why scale by $\sqrt{d_k}$?#

If query and key components are independent with zero mean and unit variance, $\mathbf{q}^\top\mathbf{k}$ has variance $d_k$. Large scores push the softmax into saturated regions where gradients vanish. Dividing by $\sqrt{d_k}$ restores unit variance.

What do attention heads learn?#

Analyses of trained models (e.g. Clark et al., 2019 on BERT; mechanistic studies of GPT-style models) found heads that:

  • attend to the previous or next token;
  • track syntactic relations (a verb attending to its direct object, determiners to their nouns);
  • resolve coreference (pronouns to their antecedents);
  • attend heavily to special tokens like [SEP] or the first token — often acting as a "no-op" or attention sink;
  • form induction heads in decoder models — circuits that look for a previous occurrence of the current token and copy what followed it, a mechanism linked to in-context learning.

Many heads are redundant; studies have pruned a large fraction of heads with little loss.

Complexity#

For sequence length $T$ and dimension $d$:

OperationTimeMemory
Projections ($\mathbf{Q}, \mathbf{K}, \mathbf{V}$)$O(Td^2)$$O(Td)$
Scores $\mathbf{Q}\mathbf{K}^\top$$O(T^2d)$$O(T^2)$ per head
Weighted sum $\mathbf{A}\mathbf{V}$$O(T^2d)$—

For short sequences the $Td^2$ projections dominate; for long sequences the quadratic $T^2$ term dominates. Doubling context length quadruples attention cost — the main obstacle to long-context models. FlashAttention computes exact attention in tiles without materialising the $T \times T$ matrix in slow GPU memory, greatly reducing memory and wall-clock time (next lectures cover efficient transformers).

Causal attention and the KV cache#

In decoder-only models, token $t$ attends only to positions $\le t$. During generation, we produce one token at a time. Recomputing keys and values for the whole prefix at each step would be wasteful, since they do not change. The KV cache stores each layer's keys and values for previous tokens; each new step computes $\mathbf{q}, \mathbf{k}, \mathbf{v}$ only for the new token and attends over the cache.

The cache grows linearly with context length, and its memory can dominate inference:

$$ \text{KV memory} = 2 \times L \times T \times n_{\text{kv}} \times d_{\text{head}} \times \text{bytes per value} \times \text{batch} $$

(factor 2 for keys and values, $L$ layers, $n_{\text{kv}}$ key/value heads).

Multi-query and grouped-query attention#

To shrink the KV cache:

  • Multi-Query Attention (MQA) (Shazeer, 2019): all query heads share a single key and value head — the cache shrinks by a factor of $h$, speeding decoding, with some quality loss.
  • Grouped-Query Attention (GQA) (Ainslie et al., 2023): query heads are divided into $g$ groups, each sharing one key/value head — a middle ground adopted by many modern LLMs.
  • Multi-head Latent Attention (MLA): compresses keys and values into a low-dimensional latent vector that is cached, further reducing memory.
python
import torch

def kv_cache_gb(layers, context, kv_heads, head_dim, batch=1, bytes_per=2):
    return 2 * layers * context * kv_heads * head_dim * batch * bytes_per / 1e9

# A hypothetical 32-layer model, head_dim 128, 32 query heads, 32K context, fp16
print("MHA (32 KV heads):", round(kv_cache_gb(32, 32_768, 32, 128), 1), "GB per sequence")
print("GQA (8 KV heads): ", round(kv_cache_gb(32, 32_768, 8, 128), 1), "GB per sequence")
print("MQA (1 KV head):  ", round(kv_cache_gb(32, 32_768, 1, 128), 2), "GB per sequence")

Cross-attention#

In encoder–decoder models (and in multimodal models that attend to image features), queries come from one sequence and keys/values from another. The same equation, a different source of $\mathbf{K}$ and $\mathbf{V}$.

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

The Transformer Architecture Explained, Block by Block

The 2017 Transformer replaced recurrence with attention and became the foundation of modern AI. We walk through embeddings, positional encoding, multi-head self-attention, feed-forward layers, residuals, normalisation, masking and the encoder–decoder design.

Intermediate⏱ 6 min#174
💬 NLP & Transformers

Positional Encodings: Sinusoidal, Learned, RoPE and ALiBi

Attention is order-blind, so transformers need positional information. We compare absolute sinusoidal and learned encodings with relative methods — rotary embeddings (RoPE) and ALiBi — and discuss extending context length.

Advanced⏱ 5 min#176
💬 NLP & Transformers

Neural Machine Translation: From Seq2Seq to Multilingual Transformers

Machine translation is one of NLP's oldest and most impactful tasks. We trace its evolution to neural systems, cover training data, subword vocabularies, back-translation, evaluation with BLEU and COMET, and the challenges of low-resource languages.

Intermediate⏱ 5 min#173