← All Posts
From Scratch · C++ · Transformers · Attention Variants · Part 4 of 13

Cross-Attention: Encoder-Decoder Pattern

Introduction

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.

LibTorch Implementation

struct CrossAttentionImpl : public torch::nn::Module { torch::nn::Linear wq{nullptr}, wk{nullptr}, wv{nullptr}; int64_t d_model, d_k; CrossAttentionImpl(int64_t d_model_) : d_model(d_model_), d_k(d_model_) { wq = register_module("wq", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wk = register_module("wk", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wv = register_module("wv", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); } torch::Tensor forward(torch::Tensor decoder_x, torch::Tensor encoder_x) { // Q from decoder auto Q = wq(decoder_x); // (batch, T_dec, d_k) // K, V from encoder auto K = wk(encoder_x); // (batch, T_enc, d_k) auto V = wv(encoder_x); // (batch, T_enc, d_k) // Scores: (batch, T_dec, d_k) @ (batch, d_k, T_enc) = (batch, T_dec, T_enc) auto scores = torch::matmul(Q, K.transpose(-2, -1)); scores = scores / sqrt((float)d_k); auto attn_weights = torch::softmax(scores, -1); // each row sums to 1 // Output: (batch, T_dec, T_enc) @ (batch, T_enc, d_k) = (batch, T_dec, d_k) auto output = torch::matmul(attn_weights, V); return output; } }; TORCH_MODULE(CrossAttention);

Shape Semantics

The key insight is that Q comes from one stream (decoder) and K, V come from another (encoder):

VariableShapeSourceMeaning
decoder_x(batch, T_dec, d_model)Decoder outputWhat the decoder has generated so far
encoder_x(batch, T_enc, d_model)Encoder outputThe 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^TAttention from each decoder to each encoder token
output(batch, T_dec, d_k)scores @ VEach 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 once auto encoder_out = encoder->forward(src_tokens); // (batch, T_src, d) // Decoder generates target sequence, one token at a time for(int i = 0; i < T_tgt; ++i) { // Decoder self-attention on tokens generated so far auto decoder_state = decoder_self_attn->forward(decoder_so_far); // Cross-attention: attend to encoder auto 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 token auto next_token = sample(next_token_logits); // Append and continue decoder_so_far = append(decoder_so_far, next_token); }

Complete Shape Trace

StepVariableShapeNotes
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 QQ(batch, T_dec, d_k)From decoder
Project KK(batch, T_enc, d_k)From encoder
Project VV(batch, T_enc, d_k)From encoder
Scoresscores(batch, T_dec, T_enc)Q @ K^T / sqrt(d_k)
Weightsattn_weights(batch, T_dec, T_enc)softmax(scores, dim=-1)
Outputoutput(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).

Previous Causal Attention