6.5 Embeddings
Your counting table stores 729 separate numbers, one per pair, and the row for “a” knows nothing about the row for “e”. What if every letter had a tiny description, just 2 numbers, and the model had to make all its predictions from those?
1Watch
An embedding is a learnable list of numbers for each symbol, stored as one row of a table called wte (word-token embeddings). To predict what follows a letter:
- Look up its row: 2 numbers.
- Score all 27 possible next symbols: a dot product with each row of a second table,
lm_head. That’slinear()from lesson 4.5. - Softmax the 27 scores into chances (6.4), and the surprise of the real next letter is the loss (5.3).
Why should a dot product make good scores? Think of each row of lm_head as an arrow for one possible next letter, and each embedding as an arrow for the current letter. Lesson 4.4 showed that a dot product is a similarity score: big when two arrows point the same way, negative when they point apart. So a letter scores high for every next letter whose arrow lines up with its own. Training turns the arrows so that pairs that often happen point the same way, and pairs that rarely happen point apart.
Gradient descent (6.2) tunes both tables. With just 2 numbers per letter the loss falls from 3.33 to 2.62. That’s not as good as the 729-number table (2.45), but it uses 108 numbers instead of 729.
Then plot each letter at its 2 numbers. Nobody told the model what a vowel is, yet a, e, i, o, u (and y) end up huddled together, because letters that are followed by similar things get similar descriptions. If “a” and “e” are both often followed by “n”, “l” and “r”, the cheapest way to score those highly for both is to point their arrows the same way, which puts them close together on the map. That sharing between similar letters is something a count table can never do.
2Explore
Nobody told the model what a vowel is. Letters that behave alike (similar next letters) end up close together, because that is the cheapest way to predict well with just 2 numbers each.
3Build
Write the forward pass, two blanks:
embed(token_id): return that letter’s row ofwte.logits_for(token_id): score every possible next letter by running the letter’s embedding throughlm_headwith thelinear()you wrote in 4.5. The input vector goes first, the table second.Show me the line
return linear(embed(token_id), lm_head)
The training loop is given, and so is its gradient formula: it was worked out by hand, and for now you can take it on trust (it gives the same slopes 6.3’s nudging would, only much faster). Module 7 shows where gradients like this come from. When your forward pass is right, training brings the loss under 2.65 and the vowels cluster.
import math, random
docs = [line.strip() for line in open('names.txt') if line.strip()]
LETTERS = '.abcdefghijklmnopqrstuvwxyz'
ID = {ch: i for i, ch in enumerate(LETTERS)}
pair_counts = {}
for name in docs:
chars = '.' + name + '.'
for a, b in zip(chars, chars[1:]):
pair_counts[(ID[a], ID[b])] = pair_counts.get((ID[a], ID[b]), 0) + 1
TOTAL = sum(pair_counts.values())
# Two tables of knobs. wte: 2 numbers per letter (its embedding).
# lm_head: 2 numbers per possible NEXT letter. Same names as in microgpt.
rng = random.Random(42)
wte = [[rng.gauss(0, 0.5) for _ in range(2)] for _ in range(27)]
lm_head = [[rng.gauss(0, 0.5) for _ in range(2)] for _ in range(27)]
def softmax(logits): # from lesson 6.4
biggest = max(logits)
exps = [math.exp(x - biggest) for x in logits]
total = sum(exps)
return [e / total for e in exps]
def linear(x, w): # microgpt line 94: one dot product per row of w
return [sum(wi * xi for wi, xi in zip(row, x)) for row in w]
# The forward pass: letter -> embedding -> 27 scores -> 27 chances.
def embed(token_id):
return [0.0, 0.0] # TODO: look up this letter's row in wte
def logits_for(token_id):
return [0.0] * 27 # TODO: linear() of the embedding with lm_head
def surprise_of_pair(a, b):
return -math.log(softmax(logits_for(a))[b])
def average_loss():
return sum(pair_counts[(a, b)] * surprise_of_pair(a, b) for (a, b) in pair_counts) / TOTAL
# Training (given): gradient descent on both tables. The gradient formula `d`
# below was worked out by hand and you can take it on trust for now. It gives
# the same numbers 6.3's nudging would, just much faster. Module 7 shows where
# gradients like this come from, and how microgpt gets them automatically.
# The right stride depends on how steep the valley is; 5.0 suits this one.
def train(epochs, lr=5.0):
for _ in range(epochs):
g_wte = [[0.0, 0.0] for _ in range(27)]
g_head = [[0.0, 0.0] for _ in range(27)]
for a in range(27):
row_total = sum(pair_counts.get((a, b), 0) for b in range(27))
if not row_total:
continue # skip to the next a at once (break would leave the loop altogether)
x, p = embed(a), softmax(logits_for(a))
for b in range(27):
d = (row_total * p[b] - pair_counts.get((a, b), 0)) / TOTAL
for j in range(2):
g_head[b][j] += d * x[j]
g_wte[a][j] += d * lm_head[b][j]
for i in range(27):
for j in range(2):
wte[i][j] -= lr * g_wte[i][j]
lm_head[i][j] -= lr * g_head[i][j]
print('before training:', round(average_loss(), 4))
train(300)
print('after 300 steps:', round(average_loss(), 4))
for ch in 'aeiou' + 'bkmt':
print(ch, [round(v, 2) for v in wte[ID[ch]]])
4Check yourself
5Unlocked in microgpt
Line 109, tok_emb = state_dict['wte'][token_id], is your embed(). Lines 143 and 144 are your final step: logits = linear(x, state_dict['lm_head']). microgpt uses 16 numbers per letter instead of 2, and everything in between (attention, the MLP: Modules 8 and 9) is what makes those 16 numbers so much smarter.
Progress is saved in this browser. Sign in to keep it across devices.