Implement Multi-Head Attention from Scratch
Company: ByteDance
Role: Applied Scientist
Category: Machine Learning
Difficulty: medium
Interview Round: Technical Screen
Implement multi-head attention, the core layer of the Transformer, from scratch. Your implementation takes query, key and value inputs and returns the attention output for every query position.
- `query` has shape `(batch, L_q, d_model)`; `key` and `value` have shape `(batch, L_k, d_model)`. For self-attention, all three are the same tensor.
- The layer has `num_heads` heads. It projects the queries, keys and values with learned matrices, runs attention in every head in parallel, concatenates the heads, and applies a learned output projection.
- The output has shape `(batch, L_q, d_model)`.
- An optional boolean mask marks which key positions each query may attend to.
Use NumPy or PyTorch tensor operations. Assume that a framework's ready-made attention layer or fused attention function is off limits, since the point is to write the layer yourself.
```hint Shapes first
Write down the shape of every tensor after every line. Most bugs in this task are a reshape or transpose that silently mixes heads with positions.
```
```hint Keep the softmax healthy
Consider how the size of the dot products changes as the per-head dimension grows, and what that does to the softmax and its gradients.
```
### Constraints and Clarifications
- `d_model` is divisible by `num_heads`; each head works in dimension `d_k = d_model / num_heads`.
- The implementation is batched, with no Python loop over heads or positions.
- Masked key positions must receive zero attention weight.
### Clarifying Questions
- A NumPy forward pass with explicit weight arguments, or a trainable PyTorch module?
- Self-attention only, or separate query and key/value inputs as in cross-attention?
- What is the mask convention: does `True` mean "may attend" or "blocked"? Is a causal mask needed?
- Should the layer also return the attention weights, and should it include dropout and bias terms?
### What a Strong Answer Covers
- Correct projections and the split into heads, with every reshape and transpose stated in shapes
- Scaled dot-product attention with a numerically stable softmax over the correct axis
- Mask semantics applied before the softmax, including a query whose keys are all masked
- Recombination of the heads and the output projection
- Time and memory complexity in sequence length and model width
- A way to test the implementation: shape checks and a comparison against a reference
### Follow-up Questions
- Add a causal mask for decoder self-attention. What changes, and what happens to a query row whose keys are all masked?
- During autoregressive generation, how would you change the layer to process one new token at a time without recomputing the keys and values of earlier tokens?
- Why use several heads of size `d_model / num_heads` instead of one head of size `d_model`? What do multi-query and grouped-query attention change?
- The sequence length grows to 32,000 tokens. Where does the layer run out of memory, and what can you do about it?
Overview: An ML coding question asking you to implement multi-head attention from scratch in NumPy or PyTorch, from the query, key and value projections to the output projection. It tests tensor shape handling, scaled dot-product attention, masking, numerical stability and complexity analysis.