Causal Multi-Head Self-Attention
Implement causal_self_attention(X, Wq, Wk, Wv, n_heads).
Xhas shape(T, D);Wq,Wk,Wveach have shape(D, D).- Project with
Q = X @ Wq(and likewise for K, V), then split theDcolumns inton_headscontiguous blocks of sizeD / n_heads. - Return the concatenated output, shape
(T, D). There is no output projection.
Softmax stability: subtract the row max before exponentiating. With -inf entries in the row this is not optional: of a large positive number overflows, and a fully-masked row would produce .
Examples
Identity projections, 1 head: full causal output
- Input
- causal_self_attention([[1, 0], [0, 1], [1, 1]], [[1, 0], [0, 1]], [[1, 0], [0, 1]], [[1, 0], [0, 1]], 1)
- Output
- [[1, 0], [0.33024, 0.66976], [0.75174, 0.75174]]
4-dim input split into 2 heads of size 2
- Input
- causal_self_attention([[1, 0, 0.5, -1], [0, 1, -0.5, 0.5], [0.5, 0.5, 1, 1], [-1, 0.5, 0, 2]], [
- Output
- [
Same input at 4 heads of size 1 differs from 2 heads of size 2
- Input
- causal_self_attention([[1, 0, 0.5, -1], [0, 1, -0.5, 0.5], [0.5, 0.5, 1, 1], [-1, 0.5, 0, 2]], [
- Output
- [
Hints
Hint 1
A triangular helper builds this mask directly. Check whether the diagonal is included.
Hint 2
Watch for this: applied the causal mask after softmax so rows no longer sum to one.
Requirements
X: (T, D) input sequenceWq: (D, D) query projectionWk: (D, D) key projectionWv: (D, D) value projectionn_heads: number of heads; D must be divisible by n_headsReturn (T, D) array. Heads are computed independently, then concatenated.
Constraints
Allowed library: NumPy only
Time limit: 200 ms, Memory: 64 MB
Where this shows up
8 employers weight this skill
4 frontier labs, 3 AI product companies, 1 enterprise vendor. Top match scores 91.
import numpy as np
def causal_self_attention(X: np.ndarray, Wq: np.ndarray, Wk: np.ndarray,
Wv: np.ndarray, n_heads: int) -> np.ndarray:
"""
Causal (masked) multi-head self-attention.
Args:
X: (T, D) input sequence
Wq: (D, D) query projection
Wk: (D, D) key projection
Wv: (D, D) value projection
n_heads: number of heads; D must be divisible by n_heads
Returns:
(T, D) array. Heads are computed independently, then concatenated.
"""
# YOUR CODE HERE
pass