Lesson 4.13 · 23 min
Cross-Attention: Connecting Encoder Output to the Decoder
When a translation model writes the French word "chats", how does it know to look back at the English word "cats" in a completely different sentence?
In short: Cross-attention is attention between two different sequences. The queries come from the sequence being generated (for example the French decoder), while the keys and values come from another source (the English encoder output, an audio clip, or a text prompt for an image model). The math is the same scaled dot-product attention; only the origin of Q versus K and V changes.
What is cross-attention?
So far we have studied self-attention, where a sequence attends to itself: every token's query is compared with the keys of tokens in the same sequence. Cross-attention connects two sequences. One sequence asks the questions (provides the queries), and the other sequence supplies the answers (provides the keys and values).
Our running example is translation. The English sentence "I love cats" goes into an encoder, a stack of layers that turns it into a set of context-rich vectors. A decoder then writes the French sentence "J'aime les chats" one token at a time. Each time the decoder is about to write a word, cross-attention lets it look over the English vectors and pick out the relevant ones.
Think of it like an interpreter with notes A human interpreter listens to a speech and takes notes (the encoder). While speaking the translation, they glance back at their notes for exactly the part they are about to say (cross-attention). Their own sentence so far decides what they look for; the notes decide what they find.
Why do we need cross-attention?
Many tasks turn one thing into another: English into French, audio into text, a text prompt into an image, a long document into a summary. The output must stay grounded in the input at every step. Self-attention alone can only look inside one sequence, so the decoder needs a separate channel into the source.
Before Transformers, encoder-decoder models built from recurrent networks compressed the whole input into a single fixed-size vector and handed that to the decoder. Long sentences lost detail because everything had to squeeze through that one bottleneck. Attention between decoder and encoder (introduced for translation around 2014–2015) fixed this by letting the decoder look at every encoder position, with weights that change at each output step. Cross-attention in the Transformer is the modern form of that idea.
Pause and think: Why does a fixed-size summary vector struggle with long inputs, while cross-attention does not?
A fixed-size vector must hold the whole input no matter how long it is, so detail gets lost. Cross-attention keeps one vector per input token and lets each output step choose which ones to read, so nothing is forced through a single bottleneck.
Query, key and value in cross-attention
The formula is unchanged: softmax(Q·Kᵀ / √dₖ)·V. What changes is where each matrix comes from:
The score matrix Q·Kᵀ therefore has shape T_tgt × T_src: one row per target token and one column per source token. It does not have to be square. A 3-word English sentence can produce a 4-token French sentence, and the weights are simply 4 × 3.
Note how "les" (the French article, which has no English counterpart here) still attends mostly to "cats", because the article is chosen to agree with the noun that follows it. Cross-attention learns these soft alignments without anyone labelling which word matches which.
Self-attention vs cross-attention
One detail in that table is easy to miss: cross-attention normally uses no causal mask. The causal mask prevents looking at future target tokens. The source sentence is fully known before decoding starts, so every target position may look at every source position. The decoder's self-attention is still causally masked.
Step-by-step working of cross-attention
Inside one decoder layer of an encoder-decoder Transformer
- Encode the source once: The encoder reads "I love cats" with bidirectional self-attention and outputs one vector per source token, H_enc.
- Masked self-attention: The decoder tokens so far ("<s> J' aime") attend to each other with a causal mask.
- Make queries: The decoder hidden states are projected with W_q into queries: "what source information do I need?"
- Make keys and values from the source: H_enc is projected with W_k and W_v. These depend only on the source, so they can be computed once and reused for every generated token.
- Attend: softmax(Q·Kᵀ / √dₖ) gives each target token a weight over source tokens; multiply by V to pull in source information.
- Feed forward: The result passes through the residual connection, normalization and the feed-forward network, then on to the next layer.
A simple example walk-through
Here is cross-attention in numpy. The decoder has produced "<s> J' aime" so far; the encoder has encoded "I love cats". Weights are random, so which word looks at which is chance. What matters is where Q, K and V come from and the shapes.
cross_attention.py
import numpy as np
np.set_printoptions(precision=2, suppress=True)
rng = np.random.default_rng(3)
src = ["I", "love", "cats"] # encoder side (English)
tgt = ["<s>", "J'", "aime"] # decoder side (French so far)
d_model, d_k = 6, 4
enc = rng.standard_normal((len(src), d_model)) # encoder output
dec = rng.standard_normal((len(tgt), d_model)) # decoder hidden states
W_q, W_k, W_v = (rng.standard_normal((d_model, d_k)) * 0.4 for _ in range(3))
def softmax(z):
z = z - z.max(axis=-1, keepdims=True)
e = np.exp(z)
return e / e.sum(axis=-1, keepdims=True)
Q = dec @ W_q # queries come from the DECODER
K = enc @ W_k # keys come from the ENCODER
V = enc @ W_v # values come from the ENCODER
print("Q", Q.shape, " K", K.shape, " V", V.shape)
weights = softmax(Q @ K.T / np.sqrt(d_k)) # (3 target x 3 source)
print("cross-attention weights (rows = French, cols = English):")
for word, row in zip(tgt, weights):
print(f" {word:>5}", row, "-> looks most at", src[row.argmax()])
out = weights @ V
print("output shape:", out.shape, "(one vector per decoder token)")
# Source length can differ from target length: add 2 more source tokens
enc5 = np.vstack([enc, rng.standard_normal((2, d_model))])
w5 = softmax(Q @ (enc5 @ W_k).T / np.sqrt(d_k))
print("with 5 source tokens, weights shape:", w5.shape)Output:
Q (3, 4) K (3, 4) V (3, 4) cross-attention weights (rows = French, cols = English): <s> [0.8 0.02 0.18] -> looks most at I J' [0.04 0.82 0.14] -> looks most at love aime [0.02 0.9 0.09] -> looks most at love output shape: (3, 4) (one vector per decoder token) with 5 source tokens, weights shape: (3, 5)
Pause and think: During generation, the decoder adds one French token per step. Which of Q, K and V must be recomputed for the new token, and which can be reused?
Only the new token's query is new. K and V come from the encoder output, which does not change during decoding, so they are computed once and cached for all steps.
Where cross-attention is used
| System | Queries from | Keys and values from |
|---|---|---|
| Original Transformer (2017), T5, BART | Decoder (output text) | Encoder (input text) |
| Whisper (speech recognition) | Text decoder | Audio encoder output |
| Stable Diffusion (text-to-image) | Image features inside the denoising U-Net | Text-encoder embeddings of the prompt |
| Flamingo-style vision-language models | Language model layers | Image features (via gated cross-attention) |
Real-world example: text-to-image When you type "a red bicycle on a beach" into a diffusion model like Stable Diffusion, the prompt is encoded into a sequence of vectors. Inside the image network, each spatial location of the image produces a query and cross-attends to the prompt vectors. That is how pixels in one region can "listen" to the word "bicycle" and pixels elsewhere to "beach".
Decoder-only chat LLMs, by contrast, usually do not have cross-attention layers. They put everything (instructions, documents, the conversation) into one sequence and rely on causal self-attention. Some multimodal models use cross-attention to inject images, while others convert images into tokens and feed them into the same sequence. Both designs exist, and which one a given product uses varies.
Why cross-attention matters, and common mistakes
- Grounding: every output step can consult the source, so outputs stay tied to the input.
- Flexible lengths: source and target lengths are independent (T_tgt × T_src weights).
- Multimodality: the source can be text, audio, image patches, or any sequence of vectors, as long as it is projected to the right width.
- Efficiency at decoding: source keys and values are computed once and reused at every step.
- Interpretability: cross-attention maps often show rough alignments, such as which source word a translated word came from.
Common mistakes Swapping the roles (taking queries from the encoder) produces outputs aligned to the source length instead of the target. Applying a causal mask to cross-attention needlessly hides parts of the source. Forgetting a padding mask on the source side lets the decoder attend to filler tokens in batched inputs. And treating cross-attention maps as exact word alignments: they are soft, and different heads and layers can disagree.
When not to use it: if all your inputs are text and you are building on a decoder-only LLM, concatenating the source into the prompt is usually simpler and works well. Cross-attention earns its place when the source is a different modality, is very long and fixed (encode once, decode many tokens), or when you want a clean separation between "what we condition on" and "what we generate".
Worked example, step by step
The lesson's code used random weights, so we could not check a row by hand. Let us do one target token with small made-up numbers. The decoder is about to write “chats”. Its query is [2, 0]. The three English tokens offer these keys and values (dₖ = 2, illustrative).
| Source token | Key | Value | Raw score with query [2, 0] |
|---|---|---|---|
| I | [0, 2] | [1, 0] | 2·0 + 0·2 = 0 |
| love | [1, 1] | [0, 1] | 2·1 + 0·1 = 2 |
| cats | [3, 0] | [1, 1] | 2·3 + 0·0 = 6 |
From scores to the output for “chats”
- Scale: Divide by √2 ≈ 1.41: [0, 2, 6] becomes [0, 1.41, 4.24].
- Softmax: e⁰ = 1, e^1.41 ≈ 4.11, e^4.24 ≈ 69.59. The sum is 74.70, so the weights are about [0.01, 0.06, 0.93]. “chats” looks almost only at “cats”.
- Blend the source values: 0.01·[1, 0] + 0.06·[0, 1] + 0.93·[1, 1] ≈ [0.94, 0.99]. The output is close to the value of “cats”.
- Now add filler: Pad the source with a fourth token. Its vector is arbitrary; say its key is [3, 0.5] and its value is [−2, 4]. Its raw score is also 6, the same as “cats”.
- See the damage: With no mask the weights become about [0.01, 0.03, 0.48, 0.48] and the output jumps to about [−0.48, 2.44]. Half of what “chats” reads is now noise.
- Fix it: Set the filler's score to −∞ before softmax. Its weight becomes 0 and we are back to [0.01, 0.06, 0.93].
Two things are worth keeping from this. First, one row of cross-attention is nothing more than a soft lookup into the source: the target token brings a question, the source brings the answers. Second, the source side needs its own mask. It is not the causal mask, which we do not use here. It is a padding mask on the source columns, and without it the result for a sentence depends on how much filler its batch happened to add.
Practice: try it yourself
We will imitate three decoding steps. The source keys and values are built once, before the loop. Each step then makes one new query and reads the source, with and without a padding mask, so we can watch both the reuse and the leak.
practice_cross_attention_steps.py
import numpy as np
np.set_printoptions(precision=2, suppress=True)
rng = np.random.default_rng(11)
src = ["I", "love", "cats", "<pad>"] # source, padded to length 4
d_model, d_k = 6, 4
enc = rng.standard_normal((len(src), d_model))
W_q, W_k, W_v = (rng.standard_normal((d_model, d_k)) * 0.6 for _ in range(3))
# Source keys and values: computed ONCE, before decoding starts
K, V = enc @ W_k, enc @ W_v
src_is_pad = np.array([s == "<pad>" for s in src])
def cross_attend(query, use_pad_mask):
scores = query @ K.T / np.sqrt(d_k) # this target token vs every source token
if use_pad_mask:
scores = np.where(src_is_pad, -np.inf, scores)
w = np.exp(scores - scores.max())
return w / w.sum()
# Decode three target tokens. Each step builds only ONE new query.
for step, word in enumerate(["J'", "aime", "les"], start=1):
query = rng.standard_normal(d_model) @ W_q # stand-in for the decoder state
leaky = cross_attend(query, use_pad_mask=False)
clean = cross_attend(query, use_pad_mask=True)
print(f"step {step} {word:>4}: no mask {leaky} masked {clean}")
print("K and V were built once and reused for", step, "steps")Output:
step 1 J': no mask [0.22 0.38 0.28 0.12] masked [0.25 0.43 0.32 0. ] step 2 aime: no mask [0.12 0.49 0.26 0.13] masked [0.13 0.56 0.3 0. ] step 3 les: no mask [0.44 0.1 0.11 0.36] masked [0.68 0.15 0.17 0. ] K and V were built once and reused for 3 steps
Now change it:
- Multiply the query by 3 (add
* 3at the end of thequery = ...line). Predict first: do the masked weights get flatter or more peaked, and does the top token change? - Add
src_is_pad[0] = Trueafter the line that buildssrc_is_pad, so “I” is hidden too. Predict the masked weights for step 1 from the current ones. - Move the line
K, V = enc @ W_k, enc @ W_vinsidecross_attend. Predict whether any printed weight changes, and count how many times the source is now projected.
Pause and think: At step 3 the unmasked row puts 36% on <pad>. Filler has no meaning. How can it earn that much weight, and what would it do to a real system?
A filler position still has a vector, so it still has a key, and a dot product with that key can be large by chance. Attention cannot know the token is meaningless unless we tell it. In a real system the output for the same sentence would then change with the amount of filler in its batch, which shows up as results that differ between batch sizes.
Pause and think: In the hand example the query for “chats” was [2, 0] and “cats” got 93%. Suppose training doubles the query to [4, 0] and changes nothing else. What happens to the weight on “cats”, and why?
It rises to almost 100%. Doubling the query doubles every raw score to [0, 4, 12], so the gaps between them double too, and softmax turns bigger gaps into a sharper split. The direction of the query decides which source token wins; its length decides by how much.
Key takeaways
- Cross-attention connects two sequences: queries from the one being generated, keys and values from the source.
- The math is identical to self-attention; the weight matrix is T_tgt × T_src and need not be square.
- It uses no causal mask, because the source is fully known; the decoder's self-attention is still masked.
- Source keys and values are computed once and reused at every decoding step.
- It powers translation (T5, BART), speech recognition (Whisper), text-to-image (Stable Diffusion) and some multimodal LLMs.
Key terms
- Cross-attention: Attention where queries come from one sequence and keys and values from another.
- Encoder: The part of a model that reads the source input and produces one vector per input token.
- Decoder: The part of a model that generates the output sequence one token at a time.
- Encoder-decoder model: A Transformer with an encoder for the input and a decoder that cross-attends to it.
- Alignment: Which source positions a target position draws from; cross-attention learns it softly.
- Conditioning: Feeding extra information (a prompt, audio, an image) that guides what a model generates.
← 4.12 Multi-Head Attention: Many Perspectives at Once · 4.14 Rotary Position Encoding: Position Without Fixed Lookup Tables →