aiengineering.guideaiengineering.guide

ML coding problem · 6 · Interview prep

Scaled dot-product attention

hardLLM internalsattentiontransformersneural-network-components

The single operation at the heart of every transformer. Implement it from scratch in pure Python — matmul, scale, row-wise softmax, matmul — and the mechanism stops being a mystery.

Problem

The problem

Implement scaled_dot_product_attention(Q, K, V) — the core of a transformer — in pure Python. Q, K, and V are matrices given as lists of lists:

  • Q is n_q x d_k (queries)
  • K is n_k x d_k (keys)
  • V is n_k x d_v (values)

Return the n_q x d_v output:

Attention(Q, K, V) = softmax( Q · Kᵀ / √d_k ) · V

The softmax is taken over each row of the scaled score matrix, so every query’s attention weights sum to 1 and the output row is a weighted blend of the rows of V. No numpy — math only. You’ll want small helpers for matrix multiply, transpose, and a stable row-wise softmax.

Concept

Attention lets each query row pull a weighted blend of value rows, where the weights come from how well that query matches each key. The formula is softmax(Q·Kᵀ / √d_k)·V. Three moves: (1) Q·Kᵀ scores every query against every key; (2) dividing by √d_k keeps the scores from growing with dimension, which would otherwise push softmax into a one-hot spike and kill gradients; (3) a row-wise softmax turns each query's scores into weights that sum to 1, and multiplying by V produces the output. Multi-head attention is just this run in parallel on projected slices, so nailing this one function is nailing the transformer.

Hints

  1. Build it in stages: scores = Q · Kᵀ, then divide every entry by sqrt(d_k), then softmax each row, then multiply the weight matrix by V.
  2. You need three helpers — matrix multiply, transpose, and a stable row softmax. zip(*M) transposes a list-of-lists.
  3. d_k is the width of Q (len(Q[0])). The output has shape (rows of Q) x (width of V), and each output row is a convex combination of the rows of V.

Solution

Try it yourself first — the reveal takes two clicks.

solution.pyPython · Pyodide
No browser Python? Run it locally instead.

The runner is a Pyodide worker (~10 MB, WebAssembly) loaded on first run. If your network or browser blocks it, copy your code and the tests below into solution.py andtest_solution.py and run pytest.

# test_solution.py
def test_output_shape():
    Q = [[1.0, 0.0], [0.0, 1.0]]
    K = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
    V = [[1.0, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]]
    out = scaled_dot_product_attention(Q, K, V)
    assert len(out) == 2
    assert all(len(row) == 3 for row in out)

def test_equal_keys_give_mean_of_values():
    # Identical (zero) keys -> uniform attention -> the mean of V's rows.
    Q = [[1.0, 1.0]]
    K = [[0.0, 0.0], [0.0, 0.0]]
    V = [[1.0, 2.0], [3.0, 4.0]]
    out = scaled_dot_product_attention(Q, K, V)
    assert abs(out[0][0] - 2.0) < 1e-9
    assert abs(out[0][1] - 3.0) < 1e-9

def test_single_key_returns_its_value():
    Q = [[1.0, 2.0]]
    K = [[1.0, 0.0]]
    V = [[5.0, 6.0]]
    out = scaled_dot_product_attention(Q, K, V)
    assert abs(out[0][0] - 5.0) < 1e-9
    assert abs(out[0][1] - 6.0) < 1e-9

def test_sharp_query_selects_one_key():
    # A query aligned strongly with key 0 should return value row 0.
    Q = [[10.0, 0.0]]
    K = [[10.0, 0.0], [0.0, 10.0]]
    V = [[1.0, 1.0], [2.0, 2.0]]
    out = scaled_dot_product_attention(Q, K, V)
    assert abs(out[0][0] - 1.0) < 1e-3
    assert abs(out[0][1] - 1.0) < 1e-3

Enter to go · Esc to close