Implement the Forward Pass of Multi-Head Attention From Scratch
Company: Meta
Role: Machine Learning Engineer
Category: Machine Learning
Difficulty: medium
Interview Round: Technical Screen
Implement the forward pass of a multi-head attention layer from scratch. The interviewer provides no test cases, so correctness has to be argued from the shapes and from small checks you write yourself.
Unless told otherwise, assume self-attention over an input `x` of shape `(batch, seq_len, d_model)`, a head count `num_heads` that divides `d_model`, learned projection matrices for queries, keys, values and the output, and an optional mask. Only the forward pass is required: no training and no backward pass.
```hint Track the shapes
Write the shape of every intermediate tensor next to the line that produces it. Most bugs in this exercise are a reshape or transpose that produces the right size but the wrong layout.
```
```hint Guard the softmax
Consider what happens to the softmax when scores are large, and when every position in a row is masked out.
```
### Clarifying Questions
- Should the implementation use NumPy, or a deep-learning framework's tensor operations without its built-in attention module?
- Is this self-attention, or cross-attention in which keys and values come from a different sequence?
- What mask format is expected: a causal mask, a padding mask or a general boolean mask, and does `True` mean "may attend" or "blocked"?
- Do the projections include bias terms, and should dropout on the attention weights be modeled?
- Should the function also return the attention weights?
### What a Strong Answer Covers
- Correct projections, head split, scaled dot-product attention per head, head concatenation and output projection, with the shape stated at every step.
- The reason for scaling by the square root of the head dimension, and a numerically stable softmax.
- Correct mask semantics, including causal masking and padding.
- Complexity in sequence length, the memory taken by the score tensor, and self-made checks such as comparison against a naive per-head loop.
### Follow-up Questions
- How would you add a key-value cache for autoregressive decoding, and how does the cost per generated token change?
- How do grouped-query and multi-query attention change the shapes and the size of the cache?
- Why is the score tensor a memory bottleneck for long sequences, and how do fused, tiled attention kernels avoid storing it?
- Where do layer normalization and the residual connection go around this layer, and why does their placement matter?
Overview: An ML coding question that asks for the forward pass of multi-head attention from scratch. It tests query, key and value projections, splitting and merging heads, scaled dot-product attention with a stable softmax, causal and padding masks, complexity in sequence length, and self-made correctness checks.