Implement Multi-Head Attention with a KV Cache and Cached Grouped-Query Attention
Company: Amazon
Role: Machine Learning Engineer
Category: Machine Learning
Difficulty: medium
Interview Round: Onsite
Implement the attention layer of a decoder-only transformer from scratch, in three steps that build on each other: standard multi-head attention, a key-value (KV) cache for token-by-token generation, and the forward pass of grouped-query attention (GQA) that uses the cache. Use a tensor library such as NumPy or PyTorch, but not a built-in attention module.
Assume the layer serves autoregressive generation, so attention is causal: each token may attend only to itself and to earlier tokens. The input is a batch of hidden states `x` of shape `(batch, seq_len, d_model)`, and the layer owns its query, key, value and output projections.
### Clarifying Questions
- NumPy or PyTorch, and may I use `einsum` and broadcasting freely?
- Are the projection weights part of the layer, or are queries, keys and values passed in already projected?
- Do all sequences in a batch have the same length, or do I need a padding mask?
- Is the cache preallocated to a maximum length, or may it grow by concatenation?
- Is positional encoding, for example rotary embeddings, applied inside this layer or outside it?
### Part 1 — Multi-head attention
Implement the forward pass of causal multi-head self-attention with `n_heads` heads of size `head_dim = d_model / n_heads`: project `x` into queries, keys and values, split them into heads, compute scaled dot-product attention with a causal mask, merge the heads, and apply the output projection.
```hint Track the shapes
Write the shape of every intermediate tensor next to the line that creates it. Most bugs in this round are a reshape or transpose that silently mixes the head axis with the time axis.
```
#### What This Part Should Cover
- Correct projections, head split and head merge, with a consistent axis order.
- Scaling by the square root of the head size, a causal mask, and a numerically stable softmax.
- Time and memory cost in terms of sequence length and model width.
### Part 2 — KV cache for incremental decoding
Extend the layer so that generation does not recompute keys and values for tokens it has already processed. The first call (prefill) receives the whole prompt; each later call receives only the newly generated token and must attend to every earlier token through a cache of their keys and values.
```hint Think in absolute positions
When only new tokens are passed in, the causal mask and any positional encoding must use each token's position in the full sequence, not its index within the current call.
```
#### Clarifying Questions for this Part
- Should one cache object serve the whole batch, and what should happen when a sequence reaches the maximum length?
#### What This Part Should Cover
- A cache layout and an update step that work for both prefill and single-token decoding.
- A correct mask when the number of queries differs from the number of keys.
- How much compute the cache saves per generated token and how much memory it costs.
- A test showing that cached decoding reproduces a full uncached forward pass.
### Part 3 — Grouped-query attention forward with the cache
The layer now has `n_heads` query heads but only `n_kv_heads` key/value heads, where `n_kv_heads` divides `n_heads`; each group of `n_heads / n_kv_heads` query heads shares one key/value head. Implement the forward pass, including the KV cache from Part 2.
```hint Decide what really needs storing
Ask which tensors must be kept per token, and whether the sharing between query heads has to be materialized at all.
```
#### What This Part Should Cover
- The mapping from query heads to shared key/value heads, implemented without a head-mixing bug.
- What the cache stores under GQA and how its size compares with multi-head attention.
- How the same code covers multi-head attention and multi-query attention as special cases.
### What a Strong Answer Covers
- One implementation that grows from Part 1 to Part 3 without rewrites, with tensor shapes documented.
- Correctness checks: cached versus uncached outputs, and GQA with `n_kv_heads = n_heads` matching multi-head attention.
- Clear reasoning about compute per generated token, cache memory, and why decoding is limited by memory bandwidth.
- Edge cases: an empty cache, cache overflow, and positions that continue across calls.
### Follow-up Questions
- How would you support a batch whose prompts have different lengths, both at prefill and during decoding?
- With rotary position embeddings, do you cache keys before or after the rotation, and why?
- How would you bound cache memory for very long conversations, for example with a sliding window, and what does that do to the mask?
- Why does GQA speed up decoding much more than it speeds up prefill?
Overview: Implement causal multi-head attention from scratch, add a key-value cache for token-by-token decoding, and extend the forward pass to grouped-query attention, where several query heads share one key and value head. Tests tensor shape discipline, masking with cached positions, and reasoning about decode compute and cache memory.