Implement Multi-Head Attention from Scratch

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

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.

|Home/Machine Learning/ByteDance
ByteDance logo
ByteDance
Oct 8, 2026
mediumApplied ScientistTechnical ScreenMachine Learning
0
0

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.

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 Guidance

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

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

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