Implement Multi-Head Attention with a KV Cache and Cached Grouped-Query Attention

Quick 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.

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.

|Home/Machine Learning/Amazon
Amazon logo
Amazon
Jan 3, 2026
mediumMachine Learning EngineerOnsiteMachine Learning
0
0

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 Guidance

  • 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.

What This Part Should Cover Guidance

  • 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.

Clarifying Questions for this Part Guidance

  • Should one cache object serve the whole batch, and what should happen when a sequence reaches the maximum length?

What This Part Should Cover Guidance

  • 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.

What This Part Should Cover Guidance

  • 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 Guidance

  • 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 Guidance

  • 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?
Loading comments...