Natural Language Processing

Attention

Attention lets a model decide which other words in a sentence matter most for the word it is currently working on.

Read these first

On this page 6
  1. Why it exists
  2. How it works
  3. The part worth reading twice
  4. Where you have already seen it
  5. Remember this
  6. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Attention is how a model decides which other words matter most for the word it is working on.

Picture yourself on a crowded railway platform, looking for one friend. Hundreds of faces pass by. Your eyes skip nearly all of them and lock onto one.

You did not study every face equally. You held a picture in your head, scanned for a match, and spent your attention where it paid off. A model does the same thing with words.

Why it exists

Read this sentence and stop at the word "it".

   The chai was hot, so it burned my tongue.

What does "it" mean? The chai. You knew instantly, and you got it from a word five positions back.

Older models struggled badly with this. They read a sentence left to right and squeezed everything they had read into one fixed-size summary. By the time they reached "it", the earlier words had been blurred together and partly lost.

It was like being told a phone number, then a shopping list, then an address. Then being asked for the phone number. Long sentences broke these models, and translation of long paragraphs was noticeably poor because of it.

Attention removed the squeeze. Every word can reach back and look directly at every other word. It picks out only the ones it needs.

How it works

Each word does three things at once. It asks a question, it advertises what it offers, and it carries content to hand over.

  Sentence:   chai      was       hot       it
                                             |
  "it" asks:  "which word am I standing for?"
                                             |
       +-------------------+-----------------+
       |         |         |
       v         v         v
     chai       was       hot
    STRONG     weak      medium      <- how well each one answers
       |                   |
       +---------+---------+
                 |
                 v
    "it" now carries mostly chai, with a bit of hot

Three ideas, with their real names:

  • The query is the question a word asks. Here, "who am I referring to?"
  • The key is what each word advertises about itself, so queries can find it.
  • The value is the content a word hands over once it has been chosen.

Think of a library. Your query is what you want. The spine of each book is its key, which is what you scan. The pages inside are the value, which is what you take away.

Every word does this for every other word, all at the same time. That is why it is called self-attention — the sentence pays attention to itself.

The part worth reading twice

The scores are not on or off. Attention is never a clean "it means chai". It is always a blend.

The word "it" ends up carrying a large share of chai. A medium share of hot. Small shares of everything else. All of it mixed into one new list of numbers.

This bothers most people the first time. You expect a decision and you get a smoothie. That blending is deliberate. It lets the model change its mind gradually as it learns. Hard choices would not allow that.

If that did not land, read this section again. It is the single idea the entire modern AI stack is built on. Almost nobody gets it on the first pass.

Where you have already seen it

  • Google Translate. Word order changes between English and Hindi. Attention is how the model works out which source word each output word should come from.
  • Every chatbot. Ask about "the second point you made" and it finds that point. Attention is the reaching-back mechanism.
  • Long document summaries. The model weighs which sentences carry the load.
  • Image generators. "A red car on a wet road" — attention is what keeps "red" attached to "car" and not to "road".

Remember this

  • Attention lets every word look directly at every other word and pick out what matters.
  • It works through three roles: a query (the question), a key (the advertisement), a value (the content).
  • The result is always a weighted blend, never a hard pick.

What to learn next

  • Transformers — the full architecture built from attention.
  • BERT — the model that made attention famous.
  • Embeddings — what the numbers being blended actually are.

Developer — Code and libraries.

Attention is four lines of NumPy. Reading those four lines carefully teaches more than any diagram.

Everything below runs on CPU in milliseconds, with hand-written numbers so you can check every value by hand.

Setup

bash
pip install numpy

Scaled dot-product attention, complete

attention.py
import numpy as np

tokens = ["chai", "was", "hot", "it"]

# Hand-made so the arithmetic stays checkable. A real model learns all three.
Q = np.array([[1, 0, 0, 0],      # chai asks about itself
              [0, 1, 0, 0],      # was  asks about itself
              [0, 0, 1, 0],      # hot  asks about itself
              [2, 0, 1, 0]],     # it   asks mostly about chai, partly about hot
             dtype=float)
K = np.eye(4)                    # each token's key is its own slot
V = np.array([[1, 0],            # chai carries "a drink"
              [0, 0],            # was  carries nothing
              [0, 1],            # hot  carries "temperature"
              [0, 0]],           # it   carries nothing by itself
             dtype=float)

d_k = K.shape[1]
scores = Q @ K.T / np.sqrt(d_k)                  # match every query against every key
weights = np.exp(scores) / np.exp(scores).sum(axis=1, keepdims=True)
output = weights @ V                             # blend the values by those weights

print(" " * 9 + "  ".join(f"{t:>5}" for t in tokens))
for tok, row in zip(tokens, weights):
    print(f"{tok:5} -> " + "  ".join(f"{w:.3f}" for w in row))

print()
print("new vector for 'it':", np.round(output[3], 3))
Output
          chai    was    hot     it
chai  -> 0.355  0.215  0.215  0.215
was   -> 0.215  0.355  0.215  0.215
hot   -> 0.215  0.215  0.355  0.215
it    -> 0.427  0.157  0.259  0.157

new vector for 'it': [0.427 0.259]

That grid is an attention map. Row it gives 0.427 to chai and 0.259 to hot, exactly as the query was built to do.

Line by line

Q @ K.T computes every query against every key in one matrix multiply. For n tokens this produces an n x n grid. Every pair, in one operation, with no loop over positions. That parallelism is the reason transformers replaced recurrent networks.

/ np.sqrt(d_k) is the "scaled" in scaled dot-product attention, and it is not decoration. The dot product of two vectors of width d_k grows with d_k. Feed large numbers into the exponential and one weight takes everything. Shown below.

The two lines computing weights are softmax: exponentiate, then divide by the row sum. Each row now sums to 1.0, which is what makes it a blend rather than an arbitrary weighting.

weights @ V performs the blend. Row it of the output is 0.427 * v_chai + 0.157 * v_was + 0.259 * v_hot + 0.157 * v_it. Since only chai and hot carry non-zero values, the result is [0.427, 0.259].

Read that output vector as a sentence. The token "it" started out meaning nothing on its own. It now carries a large amount of "drink" and a medium amount of "temperature", pulled in from its neighbours. That is contextualisation, and it is the entire function of an attention layer.

Why the scaling matters

saturation.py
import numpy as np

K = np.eye(4)
Q = np.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [2, 0, 1, 0]], dtype=float)

for label, s in [("divided by sqrt(d_k)", Q @ K.T / np.sqrt(4)),
                 ("raw, and 10x bigger ", (Q * 10) @ K.T)]:
    w = np.exp(s) / np.exp(s).sum(axis=1, keepdims=True)
    print(label, "-> weights for 'it':", np.round(w[3], 3))
Output
divided by sqrt(d_k) -> weights for 'it': [0.427 0.157 0.259 0.157]
raw, and 10x bigger  -> weights for 'it': [1. 0. 0. 0.]

The second row is a dead layer. One token takes all the weight, every other weight is effectively zero, and so is every gradient flowing through them. The model stops learning through this layer entirely.

This is the single most common cause of a transformer that trains for hours and improves not at all. Check your scaling first.

The causal mask

A model that writes text one token at a time must never see tokens that come later. Without a block, it peeks at the answer during training, scores beautifully, and then produces nonsense at generation time.

The block is a mask applied to the scores before softmax.

causal.py
import numpy as np

tokens = ["chai", "was", "hot", "it"]
Q = np.array([[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [2, 0, 1, 0]], dtype=float)
K = np.eye(4)

scores = Q @ K.T / np.sqrt(K.shape[1])
mask = np.tril(np.ones_like(scores))             # 1 on and below the diagonal
scores = np.where(mask == 1, scores, -np.inf)    # future positions become impossible

weights = np.exp(scores) / np.exp(scores).sum(axis=1, keepdims=True)

print(" " * 9 + "  ".join(f"{t:>5}" for t in tokens))
for tok, row in zip(tokens, weights):
    print(f"{tok:5} -> " + "  ".join(f"{w:.3f}" for w in row))
Output
          chai    was    hot     it
chai  -> 1.000  0.000  0.000  0.000
was   -> 0.378  0.622  0.000  0.000
hot   -> 0.274  0.274  0.452  0.000
it    -> 0.427  0.157  0.259  0.157

Everything above the diagonal is zero. chai is the first token, so it can only look at itself and gets 1.000.

Setting a score to negative infinity works because exp(-inf) is exactly 0.0 in IEEE floating point. The masked positions contribute nothing to the sum, and the remaining weights renormalise to 1.0 on their own.

Note that the it row is unchanged from the unmasked run. it is the last token, so there is no future for it to see. Masks only affect earlier positions, which is a useful thing to know when debugging.

Common mistakes

Forgetting the causal mask. Training loss drops beautifully and generation is garbage. If your validation loss looks too good to be true in a decoder, check the mask before anything else.

Forgetting to mask padding. Batches are padded to equal length. If padding is not masked, real tokens spend attention weight on empty slots, and short sequences in a batch get quietly worse than long ones. This bug is invisible in the loss curve and shows up only in per-length evaluation.

Softmax over the wrong axis. axis=1 normalises each query's row, which is correct. axis=0 normalises each column, which is meaningless and will not raise an error. It produces a model that trains slowly and never quite works.

Assuming attention weights are explanations. It is tempting to publish the attention map as a reason for the model's output. Jain & Wallace (2019), Attention is not Explanation, showed you can often construct very different attention maps that yield identical predictions. The map is a real internal quantity. It is not a faithful account of what caused the output.

Try it yourself

Change the query row for it from [2, 0, 1, 0] to [0, 0, 3, 0] and re-run. Predict the new weights before you look.

The query now points hard at slot 2, which is hot. So it should attend mostly to hot and the output vector should lean toward temperature rather than drink.

Then set that row to [0, 0, 0, 0]. A zero query gives equal scores against every key, so the weights become 0.25 each — a uniform blend that carries no information at all. That is what an untrained attention head looks like at initialisation, and it is why early training is mostly about breaking that symmetry.

What to learn next

Researcher — Mathematics and papers.

Scaled dot-product attention

text
Attention(Q, K, V) = softmax( (Q K^T) / sqrt(d_k) ) V
  • Q is the query matrix, shape n x d_k, where n is sequence length.
  • K is the key matrix, shape m x d_k. For self-attention, m = n.
  • V is the value matrix, shape m x d_v.
  • d_k is the key and query dimension per head; d_v the value dimension.
  • softmax is applied row-wise, so each query's weights sum to 1.
  • The output has shape n x d_v.

Q, K and V are linear projections of the input X of shape n x d:

text
Q = X W_Q,   K = X W_K,   V = X W_V
  • W_Q, W_K have shape d x d_k; W_V has shape d x d_v. All are learned.

Why the sqrt(d_k) denominator

Assume the components of q and k are independent with mean 0 and variance 1. Their dot product q . k = sum over i = 1..d_k of q_i k_i then has mean 0 and variance d_k, so standard deviation sqrt(d_k).

Softmax is scale-sensitive. As input magnitude grows, the distribution approaches a one-hot vector, and the Jacobian diag(p) - p p^T approaches zero — where p is the softmax output. Gradients vanish through the layer.

Dividing by sqrt(d_k) restores unit variance for the scores and keeps softmax in its responsive region. This is the argument given in footnote 4 of Vaswani et al. (2017), and the saturation.py output in the developer block is a minimal demonstration of the failure it prevents.

Multi-head attention

text
MultiHead(X) = Concat(head_1, ..., head_h) W_O
head_i = Attention(X W_Q^i, X W_K^i, X W_V^i)
  • h is the number of heads.
  • W_O has shape h*d_v x d, projecting the concatenation back to model width.
  • Convention sets d_k = d_v = d / h, so total cost matches a single full-width head.

The motivation is representational: one softmax produces one distribution per query, so a single head cannot simultaneously attend to a syntactic dependency and a coreference link. Splitting the width buys several independent distributions at no extra FLOPs.

Empirically, heads are unequal. Michel, Levy & Neubig (2019), Are Sixteen Heads Really Better than One?, pruned most heads at inference with small loss in quality. That suggests substantial redundancy at convergence.

Complexity

For sequence length n, model width d, per-layer:

OperationTimeMemory
Projections to Q, K, VO(n d^2)O(n d)
Score matrix Q K^TO(n^2 d)O(n^2) naive
Weighted sum with VO(n^2 d)O(n d)
Feed-forward blockO(n d^2)O(n d)

The crossover where attention overtakes the feed-forward block sits near n ~= d. Below it, projections dominate and long-context optimisation gains little. Above it, the n^2 term dominates everything.

Autoregressive decoding adds the KV cache. Keys and values for all previous positions are retained, costing 2 * L * n * d_head * h * bytes for L layers. At L = 32, h * d_head = 4096, n = 8192, in fp16, that is roughly 4 GB per sequence. KV cache size, not parameter count, is what limits batch size in production serving.

The efficiency literature

Exact, hardware-aware. FlashAttention (Dao et al., 2022, arXiv:2205.14135) never materialises the n x n matrix. It tiles the computation to fit in SRAM and uses online softmax with a running maximum and normaliser. Memory falls from O(n^2) to O(n) and wall-clock time improves substantially, with output identical to the naive form up to floating-point associativity. FlashAttention-2 (Dao, 2023) improves work partitioning across warps.

Cheaper KV cache. Multi-query attention (Shazeer, 2019, arXiv:1911.02150) shares one key-value head across all query heads, cutting cache size by a factor of h at some quality cost. Grouped-query attention (Ainslie et al., 2023, arXiv:2305.13245) interpolates by sharing across groups and is the current default in most open-weight models.

Sparse and windowed. Sparse Transformers (Child et al., 2019), Longformer (Beltagy et al., 2020) and BigBird combine local windows with a few global tokens. This reaches O(n * w) for window w. Quality holds for tasks with local structure and degrades where genuine long-range dependencies matter.

Linear attention. Replacing softmax(q . k) with a kernel feature map phi(q) . phi(k) allows reassociation, giving O(n d^2) (Katharopoulos et al., 2020; Choromanski et al., 2021). These consistently trail softmax attention in quality at matched compute, and the gap has not closed.

State-space alternatives. Mamba (Gu & Dao, 2023, arXiv:2312.00752) uses input-dependent selective state-space recurrence with O(n) scaling and constant-size inference state. Hybrid architectures interleaving Mamba and attention layers currently perform best in this line, which is itself informative: pure recurrence still loses on precise recall.

Position information

Attention as defined is permutation-equivariant. Shuffle the rows of X and the output shuffles identically. Position must be injected separately.

Sinusoidal encodings (Vaswani et al., 2017) add fixed vectors to the input. Rotary position embedding, RoPE (Su et al., 2021, arXiv:2104.09864), instead rotates q and k by a position-dependent angle, so their dot product depends on relative offset. RoPE is now the dominant choice, largely because context extension methods such as position interpolation and YaRN operate cleanly on its frequency basis.

Key references

  • Bahdanau, D., Cho, K. & Bengio, Y. (2014). Neural Machine Translation by Jointly Learning to Align and Translate. arXiv:1409.0473 — the original attention mechanism, additive form.
  • Luong, M.-T., Pham, H. & Manning, C. (2015). Effective Approaches to Attention-based NMT. arXiv:1508.04025 — introduces the multiplicative form.
  • Vaswani, A. et al. (2017). Attention Is All You Need. arXiv:1706.03762
  • Michel, P., Levy, O. & Neubig, G. (2019). Are Sixteen Heads Really Better than One? arXiv:1905.10650
  • Jain, S. & Wallace, B. (2019). Attention is not Explanation. arXiv:1902.10186
  • Dao, T. et al. (2022). FlashAttention. arXiv:2205.14135
  • Su, J. et al. (2021). RoFormer: Enhanced Transformer with Rotary Position Embedding. arXiv:2104.09864
  • Ainslie, J. et al. (2023). GQA. arXiv:2305.13245
  • Gu, A. & Dao, T. (2023). Mamba. arXiv:2312.00752

Current state and open problems

The n^2 cost is no longer the binding constraint it was in 2020. FlashAttention removed the memory wall, and GQA removed most of the KV cache wall. Contexts of hundreds of thousands of tokens are now routine at inference.

What has not been solved is quality across that context. Liu et al. (2023), Lost in the Middle, showed retrieval accuracy is U-shaped in position. Information at the start and end of a long context is used reliably. Information in the middle is not. Extending the window and using the window are separate achievements, and benchmarks that only test the ends will not distinguish them.

Two further open questions are worth naming precisely.

Interpretability. Attention maps are measurable but, as Jain & Wallace established, not faithful explanations. Circuit-level analysis — induction heads, path patching, sparse autoencoders on residual-stream features — is the more promising direction, and it remains labour-intensive and incomplete.

Whether softmax attention is necessary. Every cheaper alternative proposed since 2020 trails it at matched compute. Nobody has a satisfying theory of why exact softmax attention is so hard to beat. That absence of theory is the most interesting open problem in the area.

What to learn next

What to learn next

These follow on from what you just read.

  • Natural Language Processing

    BERT

    BERT reads a sentence from both directions at once and gives every word a meaning that depends on the whole sentence, which is why one pretrained BERT can be adapted to dozens of tasks cheaply.

  • Natural Language Processing

    Text classification

    Text classification sorts a piece of writing into one of a few named boxes, and the hard parts are almost never the model — they are the labels, the duplicates and the number you report.

  • Natural Language Processing

    Named entity recognition

    Named entity recognition finds the names inside a sentence — people, places, companies, dates — and the honest way to score it counts whole entities, not individual words.