ML Atlas

07 · Architectures · 5 min read · Interactive · updated

What is self-attention and how do tokens "look at" each other?

In short

In self-attention each token in a sequence forms a query, a key and a value, then gathers information from other tokens in proportion to how well they match.

What it is

Self-attention is an attention mechanism in which the queries, keys and values all come from the same sequence. Each token poses a question to every token of the same text (itself included), receives match weights from them and replaces its own representation with a weighted mixture of their values. After a single self-attention layer, every word's vector already contains information about its context.

Intuition: in the sentence "She sat on the bank and watched the river", the word "bank" is ambiguous — a financial institution or the edge of a river? Self-attention lets it "look at" "river" and "sat" and shift its representation towards the riverbank. In "The bank approved the loan" the same word gets a different representation, because it gathers information from different neighbours. This is how contextual word representations arise.

Self-attention is the heart of the transformer and of every large language model. In GPT-style models it is causal (a token sees only the tokens before it); in BERT-style models it is bidirectional (it sees the whole text).

Mechanism — why it works this way

The input is a matrix X: n tokens, each a vector of dimension d. The layer has three weight matrices: W_Q, W_K, W_V. It computes Q = X·W_Q, K = X·W_K, V = X·W_V, and then the output softmax(Q·Kᵀ / √d_k) · V. The matrix Q·Kᵀ is n×n: element (i, j) says how interested token i is in token j. The softmax works row by row, so each token distributes its "attention" — a total of 1 — across all the tokens.

Why three different projections rather than simply the product X·Xᵀ? Because the relation "who needs whom" need be neither symmetric nor based on similarity. A verb looks for its subject; a pronoun, for the noun it refers to. Separate W_Q and W_K make it possible to learn that token A asks for a feature that token B offers, even though A and B themselves are not similar. W_V separates "how to find me" from "what I pass on".

Self-attention differs from RNNs and convolutions in two important ways. First, every token has direct access to every other token — the path between the first and the thousandth word has length 1, not 1,000 recurrent steps. Second, all positions are computed in parallel with a single matrix multiplication, which suits GPUs perfectly. These are the two main reasons transformers replaced RNNs.

The price is quadratic cost: the n×n matrix has to be computed in every layer and every head. Doubling the context length means four times as many pairs. That is why so much research goes into efficient attention variants, and why models' context windows are limited.

An important, less obvious property: self-attention does not know the order. If you permute the input tokens, the outputs are permuted in exactly the same way, but their values do not change (this is called permutation equivariance). To self-attention, "dog bites man" and "man bites dog" are the same set of words. Order information has to be added separately — through positional encoding.

By example

Three tokens A, B, C with vectors X_A = [1, 0], X_B = [0, 1], X_C = [1, 1]. Projections: W_Q with rows [1, 0] and [0, 2], W_K with rows [0, 1] and [2, 0], W_V = the identity matrix. This gives Q_A = [1, 0], Q_B = [0, 2], Q_C = [1, 2] and K_A = [0, 1], K_B = [2, 0], K_C = [2, 1]. The scores Q·Kᵀ / √2 in row C are [1.41, 1.41, 2.83], so after the softmax token C splits its attention: 0.16 on A, 0.16 on B and 0.67 on itself. Token A has weights [0.11, 0.45, 0.45] — it hardly looks at itself, only at B and C. The outputs (weight matrix times V): A → [0.55, 0.89], B → [0.89, 0.55], C → [0.84, 0.84]. With the order C, A, B the output is exactly the same three vectors in the new order — we checked this numerically.

The cost in numbers: with 1,024 tokens the attention matrix has 1,048,576 entries per head and layer; with 8,192 tokens, 67,108,864, i.e. about 134 MB in 16-bit format if held in memory in full. The number of weights in the layer does not depend on the length of the text: for d = 768 (as in GPT-2 small) the four projection matrices (Q, K, V and output) with biases amount to 4·768² + 4·768 = 2,362,368 parameters.

In practice

  • PyTorch: nn.MultiheadAttention(d, num_heads, batch_first=True) called as attn(x, x, x) — the same input three times is exactly self-attention.
  • For generative models: F.scaled_dot_product_attention(q, k, v, is_causal=True) or an explicit triangular mask.
  • Remember positional encoding — without it the model is blind to word order.
  • For long sequences, efficient implementations (FlashAttention) are used that never materialize the n×n matrix in GPU memory, although the number of operations still grows quadratically.
  • Common mistake: confusing the padding mask with the causal mask, or inverting the convention (True = "block" for boolean masks in PyTorch's nn.MultiheadAttention).

Frequently asked questions

How is self-attention different from ordinary attention?
In classic attention the queries come from one sequence and the keys and values from another (e.g. a translation looking at the source sentence). In self-attention all three come from the same text, so tokens build their context from one another.
Why does the cost of self-attention grow quadratically?
Because each of the n tokens computes its match with each of the n tokens, which gives n² pairs. With a text ten times longer there are a hundred times as many pairs.
Does a token also look at itself?
Yes, the diagonal of the attention matrix is each token's attention to itself. It often carries substantial weight, because a token's own content matters to it, but the model can learn to look mainly at others.

Sources

  • Vaswani et al. "Attention Is All You Need", NeurIPS 2017, arXiv:1706.03762.
  • Devlin, Chang, Lee, Toutanova "BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding", NAACL 2019.
  • Zhang et al. "Dive into Deep Learning", d2l.ai, ch. 11.6 ("Self-Attention and Positional Encoding").
  • Dao et al. "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness", NeurIPS 2022.
  • PyTorch documentation: torch.nn.MultiheadAttention, https://pytorch.org/docs/stable/generated/torch.nn.MultiheadAttention.html

See also