PIXELBANKv9.1.0
Menu

Implement multi-head attention by splitting Q, K, V into multiple heads.

Given Q, K, V of shape (n, d) and number of heads h:

  1. Split each into h heads: reshape (n, d) → (n, h, d/h) → (h, n, d/h)
  2. Apply scaled dot-product attention per head
  3. Concatenate heads: (h, n, d/h) → (n, d)

Assume d is divisible by h. No linear projections needed — just split and concat.

Input:

  • Line 1: n d h
  • Next n lines: Q matrix
  • Next n lines: K matrix
  • Next n lines: V matrix

Output: Multi-head attention output (n, d), values rounded to 4 decimal places.

Example:

Input:
2 4 2
1 0 0 1
0 1 1 0
1 0 0 1
0 1 1 0
1 2 3 4
5 6 7 8
Output:
[2.3210 3.3210 4.3210 5.3210]
[3.6790 4.6790 5.6790 6.6790]
Reasoning:
  • Split QQ, KK, VV column-wise into h=2h = 2 heads of d/h=2d/h = 2 features (head 1 = columns 0–1, head 2 = columns 2–3): Q1=K1=[1001]Q_1 = K_1 = \begin{bmatrix}1&0\\0&1\end{bmatrix}, Q2=K2=[0110]Q_2 = K_2 = \begin{bmatrix}0&1\\1&0\end{bmatrix}, V1=[1256]V_1 = \begin{bmatrix}1&2\\5&6\end{bmatrix}, V2=[3478]V_2 = \begin{bmatrix}3&4\\7&8\end{bmatrix}.
  • Head 1 scores: Q1K1T/2=[0.7071000.7071]Q_1K_1^T/\sqrt{2} = \begin{bmatrix}0.7071&0\\0&0.7071\end{bmatrix}. Row-wise softmax gives e0.7071e0.7071+e0=2.02813.0281=0.6698\frac{e^{0.7071}}{e^{0.7071}+e^{0}} = \frac{2.0281}{3.0281} = 0.6698, so the weights are [0.66980.33020.33020.6698]\begin{bmatrix}0.6698&0.3302\\0.3302&0.6698\end{bmatrix}.
  • Head 2: Q2K2TQ_2K_2^T is also the identity, so it gets the same weights.
  • Weighted values: head 1 row 0 is 0.6698â‹…[1,2]+0.3302â‹…[5,6]=[2.3210,3.3210]0.6698\cdot[1,2] + 0.3302\cdot[5,6] = [2.3210, 3.3210] and head 2 row 0 is 0.6698â‹…[3,4]+0.3302â‹…[7,8]=[4.3210,5.3210]0.6698\cdot[3,4] + 0.3302\cdot[7,8] = [4.3210, 5.3210]. Row 1 gives [3.6790,4.6790][3.6790, 4.6790] and [5.6790,6.6790][5.6790, 6.6790].
  • Concatenate the two heads for each token and round to 4 decimals: [2.32103.32104.32105.32103.67904.67905.67906.6790]\begin{bmatrix}2.3210&3.3210&4.3210&5.3210\\3.6790&4.6790&5.6790&6.6790\end{bmatrix}.

Constraints:

  • d is divisible by h
  • 1 <= h <= d, 1 <= n <= 10
  • No projection matrices — just split/concat
  • Round to 4 decimal places
solution.py

Test Results

0/0
Run code to see test results.