ML coding problem · 6 · Interview prep
Scaled dot-product attention
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:
Qisn_q x d_k(queries)Kisn_k x d_k(keys)Visn_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
- 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.
- You need three helpers — matrix multiply, transpose, and a stable row softmax. zip(*M) transposes a list-of-lists.
- 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.
import math
def _matmul(A, B):
inner = len(B)
cols = len(B[0])
return [
[sum(A[i][k] * B[k][j] for k in range(inner)) for j in range(cols)]
for i in range(len(A))
]
def _transpose(M):
return [list(col) for col in zip(*M)]
def _softmax_row(row):
m = max(row)
exps = [math.exp(x - m) for x in row]
total = sum(exps)
return [e / total for e in exps]
def scaled_dot_product_attention(Q, K, V):
d_k = len(Q[0])
scale = math.sqrt(d_k)
scores = _matmul(Q, _transpose(K))
scores = [[s / scale for s in row] for row in scores]
weights = [_softmax_row(row) for row in scores]
return _matmul(weights, V)
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