Skip to main content
Course map
Module 6 · Learning = walking downhill · train1.py · ~25 min

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:

  1. Look up its row: 2 numbers.
  2. Score all 27 possible next symbols: a dot product with each row of a second table, lm_head. That’s linear() from lesson 4.5.
  3. 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

·abcdefghijklmnopqrstuvwxyz
Loss during training (300 steps)
3.332.62

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:

  1. embed(token_id): return that letter’s row of wte.
  2. logits_for(token_id): score every possible next letter by running the letter’s embedding through lm_head with the linear() you wrote in 4.5. The input vector goes first, the table second.
    Show me the linereturn 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

1. An embedding is…
2. How many knobs does the 2-number embedding model have?
3. Predict: before training (small random knobs), what does average_loss() print, roughly?
4. Why do the vowels end up near each other?

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.

microgpt.py71 / 175 lines learned
1"""
2The most atomic way to train and run inference for a GPT in pure, dependency-free Python.
3This file is the complete algorithm.
4Everything else is just efficiency.
5
6@karpathy
7"""
8
9import os # os.path.exists
10import math # math.log, math.exp
11import random # random.seed, random.choices, random.gauss, random.shuffle
12random.seed(42) # Let there be order among chaos
13
14# Let there be a Dataset `docs`: list[str] of documents (e.g. a list of names)
15if not os.path.exists('input.txt'):
16 import urllib.request
17 names_url = 'https://raw.githubusercontent.com/karpathy/makemore/988aa59/names.txt'
18 urllib.request.urlretrieve(names_url, 'input.txt')
19docs = [line.strip() for line in open('input.txt') if line.strip()]
20random.shuffle(docs)
21print(f"num docs: {len(docs)}")
22
23# Let there be a Tokenizer to translate strings to sequences of integers ("tokens") and back
24uchars = sorted(set(''.join(docs))) # unique characters in the dataset become token ids 0..n-1
25BOS = len(uchars) # token id for a special Beginning of Sequence (BOS) token
26vocab_size = len(uchars) + 1 # total number of unique tokens, +1 is for BOS
27print(f"vocab size: {vocab_size}")
28
29# Let there be Autograd to recursively apply the chain rule through a computation graph
30class Value:
31 __slots__ = ('data', 'grad', '_children', '_local_grads') # Python optimization for memory usage
32
33 def __init__(self, data, children=(), local_grads=()):
34 self.data = data # scalar value of this node calculated during forward pass
35 self.grad = 0 # derivative of the loss w.r.t. this node, calculated in backward pass
36 self._children = children # children of this node in the computation graph
37 self._local_grads = local_grads # local derivative of this node w.r.t. its children
38
39 def __add__(self, other):
40 other = other if isinstance(other, Value) else Value(other)
41 return Value(self.data + other.data, (self, other), (1, 1))
42
43 def __mul__(self, other):
44 other = other if isinstance(other, Value) else Value(other)
45 return Value(self.data * other.data, (self, other), (other.data, self.data))
46
47 def __pow__(self, other): return Value(self.data**other, (self,), (other * self.data**(other-1),))
48 def log(self): return Value(math.log(self.data), (self,), (1/self.data,))
49 def exp(self): return Value(math.exp(self.data), (self,), (math.exp(self.data),))
50 def relu(self): return Value(max(0, self.data), (self,), (float(self.data > 0),))
51 def __neg__(self): return self * -1
52 def __radd__(self, other): return self + other
53 def __sub__(self, other): return self + (-other)
54 def __rsub__(self, other): return other + (-self)
55 def __rmul__(self, other): return self * other
56 def __truediv__(self, other): return self * other**-1
57 def __rtruediv__(self, other): return other * self**-1
58
59 def backward(self):
60 topo = []
61 visited = set()
62 def build_topo(v):
63 if v not in visited:
64 visited.add(v)
65 for child in v._children:
66 build_topo(child)
67 topo.append(v)
68 build_topo(self)
69 self.grad = 1
70 for v in reversed(topo):
71 for child, local_grad in zip(v._children, v._local_grads):
72 child.grad += local_grad * v.grad
73
74# Initialize the parameters, to store the knowledge of the model
75n_layer = 1 # depth of the transformer neural network (number of layers)
76n_embd = 16 # width of the network (embedding dimension)
77block_size = 16 # maximum context length of the attention window (note: the longest name is 15 characters)
78n_head = 4 # number of attention heads
79head_dim = n_embd // n_head # derived dimension of each head
80matrix = lambda nout, nin, std=0.08: [[Value(random.gauss(0, std)) for _ in range(nin)] for _ in range(nout)]
81state_dict = {'wte': matrix(vocab_size, n_embd), 'wpe': matrix(block_size, n_embd), 'lm_head': matrix(vocab_size, n_embd)}
82for i in range(n_layer):
83 state_dict[f'layer{i}.attn_wq'] = matrix(n_embd, n_embd)
84 state_dict[f'layer{i}.attn_wk'] = matrix(n_embd, n_embd)
85 state_dict[f'layer{i}.attn_wv'] = matrix(n_embd, n_embd)
86 state_dict[f'layer{i}.attn_wo'] = matrix(n_embd, n_embd)
87 state_dict[f'layer{i}.mlp_fc1'] = matrix(4 * n_embd, n_embd)
88 state_dict[f'layer{i}.mlp_fc2'] = matrix(n_embd, 4 * n_embd)
89params = [p for mat in state_dict.values() for row in mat for p in row] # flatten params into a single list[Value]
90print(f"num params: {len(params)}")
91
92# Define the model architecture: a function mapping tokens and parameters to logits over what comes next
93# Follow GPT-2, blessed among the GPTs, with minor differences: layernorm -> rmsnorm, no biases, GeLU -> ReLU
94def linear(x, w):
95 return [sum(wi * xi for wi, xi in zip(wo, x)) for wo in w]
96
97def softmax(logits):
98 max_val = max(val.data for val in logits)
99 exps = [(val - max_val).exp() for val in logits]
100 total = sum(exps)
101 return [e / total for e in exps]
102
103def rmsnorm(x):
104 ms = sum(xi * xi for xi in x) / len(x)
105 scale = (ms + 1e-5) ** -0.5
106 return [xi * scale for xi in x]
107
108def gpt(token_id, pos_id, keys, values):
109 tok_emb = state_dict['wte'][token_id] # token embedding
110 pos_emb = state_dict['wpe'][pos_id] # position embedding
111 x = [t + p for t, p in zip(tok_emb, pos_emb)] # joint token and position embedding
112 x = rmsnorm(x) # note: not redundant due to backward pass via the residual connection
113
114 for li in range(n_layer):
115 # 1) Multi-head Attention block
116 x_residual = x
117 x = rmsnorm(x)
118 q = linear(x, state_dict[f'layer{li}.attn_wq'])
119 k = linear(x, state_dict[f'layer{li}.attn_wk'])
120 v = linear(x, state_dict[f'layer{li}.attn_wv'])
121 keys[li].append(k)
122 values[li].append(v)
123 x_attn = []
124 for h in range(n_head):
125 hs = h * head_dim
126 q_h = q[hs:hs+head_dim]
127 k_h = [ki[hs:hs+head_dim] for ki in keys[li]]
128 v_h = [vi[hs:hs+head_dim] for vi in values[li]]
129 attn_logits = [sum(q_h[j] * k_h[t][j] for j in range(head_dim)) / head_dim**0.5 for t in range(len(k_h))]
130 attn_weights = softmax(attn_logits)
131 head_out = [sum(attn_weights[t] * v_h[t][j] for t in range(len(v_h))) for j in range(head_dim)]
132 x_attn.extend(head_out)
133 x = linear(x_attn, state_dict[f'layer{li}.attn_wo'])
134 x = [a + b for a, b in zip(x, x_residual)]
135 # 2) MLP block
136 x_residual = x
137 x = rmsnorm(x)
138 x = linear(x, state_dict[f'layer{li}.mlp_fc1'])
139 x = [xi.relu() for xi in x]
140 x = linear(x, state_dict[f'layer{li}.mlp_fc2'])
141 x = [a + b for a, b in zip(x, x_residual)]
142
143 logits = linear(x, state_dict['lm_head'])
144 return logits
145
146# Let there be Adam, the blessed optimizer and its buffers
147learning_rate, beta1, beta2, eps_adam = 0.01, 0.85, 0.99, 1e-8
148m = [0.0] * len(params) # first moment buffer
149v = [0.0] * len(params) # second moment buffer
150
151# Repeat in sequence
152num_steps = 1000 # number of training steps
153for step in range(num_steps):
154
155 # Take single document, tokenize it, surround it with BOS special token on both sides
156 doc = docs[step % len(docs)]
157 tokens = [BOS] + [uchars.index(ch) for ch in doc] + [BOS]
158 n = min(block_size, len(tokens) - 1)
159
160 # Forward the token sequence through the model, building up the computation graph all the way to the loss
161 keys, values = [[] for _ in range(n_layer)], [[] for _ in range(n_layer)]
162 losses = []
163 for pos_id in range(n):
164 token_id, target_id = tokens[pos_id], tokens[pos_id + 1]
165 logits = gpt(token_id, pos_id, keys, values)
166 probs = softmax(logits)
167 loss_t = -probs[target_id].log()
168 losses.append(loss_t)
169 loss = (1 / n) * sum(losses) # final average loss over the document sequence. May yours be low.
170
171 # Backward the loss, calculating the gradients with respect to all model parameters
172 loss.backward()
173
174 # Adam optimizer update: update the model parameters based on the corresponding gradients
175 lr_t = learning_rate * (1 - step / num_steps) # linear learning rate decay
176 for i, p in enumerate(params):
177 m[i] = beta1 * m[i] + (1 - beta1) * p.grad
178 v[i] = beta2 * v[i] + (1 - beta2) * p.grad ** 2
179 m_hat = m[i] / (1 - beta1 ** (step + 1))
180 v_hat = v[i] / (1 - beta2 ** (step + 1))
181 p.data -= lr_t * m_hat / (v_hat ** 0.5 + eps_adam)
182 p.grad = 0
183
184 print(f"step {step+1:4d} / {num_steps:4d} | loss {loss.data:.4f}", end='\r')
185
186# Inference: may the model babble back to us
187temperature = 0.5 # in (0, 1], control the "creativity" of generated text, low to high
188print("\n--- inference (new, hallucinated names) ---")
189for sample_idx in range(20):
190 keys, values = [[] for _ in range(n_layer)], [[] for _ in range(n_layer)]
191 token_id = BOS
192 sample = []
193 for pos_id in range(block_size):
194 logits = gpt(token_id, pos_id, keys, values)
195 probs = softmax([l / temperature for l in logits])
196 token_id = random.choices(range(vocab_size), weights=[p.data for p in probs])[0]
197 if token_id == BOS:
198 break
199 sample.append(uchars[token_id])
200 print(f"sample {sample_idx+1:2d}: {''.join(sample)}")
this lessonlearnedKarpathy’s original
Was this lesson clear?
End of Module 6: check what you learnedNext available lesson → 7.1 Computation graphs