In self-attention, queries, keys, and values come from the same source. But in encoder-decoder architectures (machine translation, summarization, image captioning), the decoder needs to attend to the encoder's output, not to itself. This is cross-attention: queries come from the decoder, while keys and values come from the encoder.
This post implements a C++ cross-attention layer where decoder and encoder sequences can have different lengths, and shows the two-stream data flow.
The Cross-Attention Principle
In a translation task:
Encoder input: "Hello world" → encoder output: (batch, T_enc, d)
Decoder input: "Hola" (so far) → decoder state: (batch, T_dec, d)
Cross-attention:
Q comes from decoder: shape (batch, T_dec, d_k)
K, V come from encoder: shape (batch, T_enc, d_k)
Result: (batch, T_dec, d_k)
Each decoder token attends to all encoder tokens.
The key difference: $Q$ and $K$ may have different sequence lengths. The score matrix is $(T_{\text{dec}} \times T_{\text{enc}})$ instead of square.
The key insight is that Q comes from one stream (decoder) and K, V come from another (encoder):
Variable
Shape
Source
Meaning
decoder_x
(batch, T_dec, d_model)
Decoder output
What the decoder has generated so far
encoder_x
(batch, T_enc, d_model)
Encoder output
The encoder's full representation of input
Q
(batch, T_dec, d_k)
Q = wq(decoder_x)
What each decoder position is looking for
K
(batch, T_enc, d_k)
K = wk(encoder_x)
What each encoder position represents
V
(batch, T_enc, d_k)
V = wv(encoder_x)
What each encoder position contributes
scores
(batch, T_dec, T_enc)
Q @ K^T
Attention from each decoder to each encoder token
output
(batch, T_dec, d_k)
scores @ V
Each decoder position has combined encoder information
Interactive Data Flow
▶ Two-Stream Cross-Attention Flow
Watch how the decoder query stream and encoder context stream flow through cross-attention.
Usage in Encoder-Decoder Models
// Encoder processes input sequence onceauto encoder_out = encoder->forward(src_tokens); // (batch, T_src, d)// Decoder generates target sequence, one token at a timefor(int i = 0; i < T_tgt; ++i) {
// Decoder self-attention on tokens generated so farauto decoder_state = decoder_self_attn->forward(decoder_so_far);
// Cross-attention: attend to encoderauto cross_out = cross_attn->forward(decoder_state, encoder_out);
// Feed-forward, layer norm, etc.auto next_token_logits = ff->forward(cross_out);
// Sample or argmax to get next tokenauto next_token = sample(next_token_logits);
// Append and continue
decoder_so_far = append(decoder_so_far, next_token);
}
Complete Shape Trace
Step
Variable
Shape
Notes
Input (decoder)
decoder_x
(batch, T_dec, d_model)
Decoder output so far
Input (encoder)
encoder_x
(batch, T_enc, d_model)
Encoder output (fixed)
Project Q
Q
(batch, T_dec, d_k)
From decoder
Project K
K
(batch, T_enc, d_k)
From encoder
Project V
V
(batch, T_enc, d_k)
From encoder
Scores
scores
(batch, T_dec, T_enc)
Q @ K^T / sqrt(d_k)
Weights
attn_weights
(batch, T_dec, T_enc)
softmax(scores, dim=-1)
Output
output
(batch, T_dec, d_k)
weights @ V
Key Insights
Cross-attention joins two independent sequences. Unlike self-attention's single input, cross-attention takes two: one for Q and one for K, V. This enables encoder-decoder and other multi-stream architectures.
Sequences can have different lengths. The decoder and encoder may process different-length inputs. Cross-attention naturally handles this via a non-square score matrix $(T_{\text{dec}} \times T_{\text{enc}})$.
Encoder is processed once, decoder incrementally. In typical generation, the encoder runs once on the full input. The decoder then runs iteratively (or in parallel during training), attending to the fixed encoder output at each step.
What Comes Next
Self-attention and cross-attention attend to all positions. But for long sequences, this $O(T^2)$ cost is prohibitive. The next posts explore structured sparsity patterns: Sliding-Window Attention (attend to a band around each position) and Global Attention (combine sliding-window with a few "global" tokens).