8.3 Scaled dot-product attention
The counting model saw one letter back. The embedding model saw one letter back. Today a letter gets to look at every letter before it and decide for itself which ones matter. This is the idea that made GPT possible.
1Watch
One head of attention, for the newest letter, takes three steps:
- Score its query against every key so far, its own included (scaled dot product, 8.2).
- Softmax the scores (6.4). Now they’re attention weights: positive, adding up to 100%.
- Blend the values, each weighted by its attention. A letter with 60% attention contributes 60% of the result.
A worked blend
In the lab’s tiny example the weights come out 0.576, 0.14 and 0.284, and the three values are [10, 0], [0, 10] and [5, 5]. The blend is 0.576 × [10, 0] + 0.14 × [0, 10] + 0.284 × [5, 5], done one number at a time:
- first number: 0.576 × 10 + 0.14 × 0 + 0.284 × 5 = 5.76 + 0 + 1.42 = 7.18
- second number: 0.576 × 0 + 0.14 × 10 + 0.284 × 5 = 0 + 1.4 + 1.42 = 2.82
So the output is [7.18, 2.82]: mostly the first value, because its key matched best. Like mixing paint in those proportions.
What the blend is for
Here is the whole trip one letter makes through microgpt: embed (8.1) → normalise (8.5) → attention → add back (8.6) → MLP (9.2) → lm_head turns the vector into 27 scores, one per possible next letter (6.5). Attention’s job is to pack what the earlier letters say into this letter’s vector before that guess. Without it, the guess after “emm” could only see the last “m”.
Nothing about “which letters matter” is hand-coded: the query, key and value matrices are knobs, and training tunes them.
The heatmap shows the real trained microgpt reading a name, one head at a time (each head is its own search on its own 4 numbers; 9.1 shows all four side by side). Honest note: in a model this small (1 layer, see 9.3; 1,000 steps), the heads are fuzzy rather than neatly specialised. Each spreads its attention over recent letters, the start marker and itself. Bigger models grow much sharper patterns.
2Explore
Try it: with “kaylee” and head 2, look at the last row. Where does the final “e” look most? (Answer: the “l”, 43%.) Tap a square to read its percentage. Then type your own name.
3Build
Two blanks in attend(q, keys, values):
- Softmax the scores into weights:
softmax(scores). - Blend: one sum per output position
j, exactly like the worked example. The shape is[sum(… for w, v in zip(weights, values)) for j in range(len(values[0]))], and the … usesv[j]. (You can’t multiply a whole list by 0.576 in Python, sow * von its own fails.)
import json, math
# The real microgpt after 1,000 training steps (same weights the widgets use).
M = json.load(open("microgpt-trained.json"))
W, UCHARS = M["weights"], M["uchars"]
BOS = len(UCHARS)
def linear(x, w): # microgpt line 94
return [sum(wi * xi for wi, xi in zip(row, x)) for row in w]
def softmax(z): # lesson 6.4
m = max(z)
e = [math.exp(v - m) for v in z]
t = sum(e)
return [v / t for v in e]
# One head of attention, for the newest token:
# 1. score every key against my query (scaled dot product)
# 2. softmax the scores into attention weights (they add up to 1)
# 3. blend the values, each weighted by its attention
def attend(q, keys, values):
scores = [sum(qi * ki for qi, ki in zip(q, k)) / math.sqrt(len(q)) for k in keys]
weights = [1 / len(keys)] * len(keys) # TODO: softmax the scores instead
blended = values[-1] # TODO: for each position j, sum weight * value[j] over all values
return blended, weights
# A tiny hand-made example: the query looks for "vowel-ish" keys.
q = [1.0, 0.0]
keys = [[2.0, 0.0], [0.0, 2.0], [1.0, 1.0]] # vowel, consonant, a bit of both
values = [[10.0, 0.0], [0.0, 10.0], [5.0, 5.0]]
# try runs the lines under it. If one fails with a TypeError, Python jumps to
# the matching except block and prints a hint instead of stopping with an error.
try:
out, w = attend(q, keys, values)
print('attention weights:', [round(x, 3) for x in w])
print('blended value: ', [round(x, 3) for x in out])
except TypeError as err:
print('TypeError:', err)
print('Hint: Python cannot multiply a whole list by a number (0.576 * [10, 0] fails).')
print('Blend one output position j at a time: one sum per j, using v[j].')
except IndexError as err:
print('IndexError:', err)
print('Hint: j counts the numbers inside ONE value, so use range(len(values[0])).')
4Check yourself
5Unlocked in microgpt
Line 123 starts collecting the head outputs, and lines 129 to 131 are your attend(): scores, softmax, blended values. Lines 124 to 128 split the work across 4 heads, which is lesson 9.1.
Progress is saved in this browser. Sign in to keep it across devices.