Causal Multi-Head Self-Attention

~40 minimplementation

Implement causal_self_attention(X, Wq, Wk, Wv, n_heads).

  • X has shape (T, D); Wq, Wk, Wv each have shape (D, D).
  • Project with Q = X @ Wq (and likewise for K, V), then split the D columns into n_heads contiguous blocks of size D / 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 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

  • Return (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

~40 min

8 employers weight this skill

4 frontier labs, 3 AI product companies, 1 enterprise vendor. Top match scores 91.

Python
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
Loading docs…

The AI Mentor needs an account

It reads your code and the failing tests and nudges you toward the fix without handing you the answer. Free accounts get it on every problem you're working on today.