9.3 Stacking layers
You have every part: embeddings, attention, the MLP, norms and residual highways. Today you snap them together into gpt(), and it produces exactly the same predictions as Karpathy’s trained model.
1Watch
gpt(token, position, keys, values) takes one token and returns 27 scores for what comes next:
- token embedding + position embedding, then RMSNorm (8.1, 8.5)
- attention block with its residual (8.2 to 9.1)
- MLP block with its residual (9.2)
lm_head: 16 numbers to 27 scores (6.5)
Steps 2 and 3 form one layer, also called a Transformer block. microgpt has n_layer = 1. Bigger GPTs stack dozens of identical layers, each with its own knobs, by looping over steps 2 and 3. That’s why microgpt’s code has for li in range(n_layer), with table names like f'layer{li}.attn_wq' and one key cache per layer, keys[li].
One more size: block_size = 16 is the context window, the most positions the model can read at once. It is why wpe has 16 rows (8.1), and why training trims each name to fit.
Residuals on lists. In 8.6 you wrote x = x + layer(x) for single numbers. Here x is a list of 16 numbers, and on lists + glues them end to end (32 numbers). Add them number by number instead:
x = [a + b for a, b in zip(x, x_residual)]2Explore
3Build
multi_head, mlp and rmsnorm are given. You write the body of gpt() after the embeddings: the layer loop for li in range(N_LAYER) with its two blocks and their residuals, using keys[li] and names like W[f'layer{li}.attn_wq'], then lm_head. The comments in the starter list the steps. This is your rehearsal for the capstone. When it’s right, your gpt() matches microgpt’s predictions letter by letter. If you glue lists by mistake, linear stops with a message saying so.
import json, math
# The real microgpt after 1,000 training steps.
M = json.load(open("microgpt-trained.json"))
W, UCHARS = M["weights"], M["uchars"]
BOS, N_LAYER, N_HEAD, N_EMBD = len(UCHARS), 1, 4, 16
HEAD_DIM = N_EMBD // N_HEAD
def linear(x, w):
assert len(x) == len(w[0]), (f'linear got {len(x)} numbers, but this table expects {len(w[0])}. '
'Did you write x + x_residual? On lists, + glues them end to end; '
'add number by number: [a + b for a, b in zip(x, x_residual)]')
return [sum(wi * xi for wi, xi in zip(row, x)) for row in w]
def softmax(z):
m = max(z)
e = [math.exp(v - m) for v in z]
t = sum(e)
return [v / t for v in e]
def rmsnorm(x):
ms = sum(v * v for v in x) / len(x)
return [v * (ms + 1e-5) ** -0.5 for v in x]
def multi_head(q, keys, values): # lesson 9.1; keys, values: one layer's cache
assert keys and not isinstance(keys[0][0], list), (
"multi_head needs this layer's list of keys, with this token's key already in it: "
'append to keys[li], then pass keys[li]')
out = []
for h in range(N_HEAD):
s = h * HEAD_DIM
w = softmax([sum(q[s + j] * k[s + j] for j in range(HEAD_DIM)) / HEAD_DIM ** 0.5 for k in keys])
out.extend(sum(w[t] * values[t][s + j] for t in range(len(values))) for j in range(HEAD_DIM))
return out
def mlp(x, li): # lesson 9.2, for layer li
h = [max(0.0, v) for v in linear(x, W[f'layer{li}.mlp_fc1'])]
return linear(h, W[f'layer{li}.mlp_fc2'])
# The whole model, for one token at one position (microgpt's gpt()).
# keys and values hold one list per layer: keys[li] is layer li's cache.
# Table names carry the layer number: W[f'layer{li}.attn_wq'] and so on.
def gpt(token_id, pos_id, keys, values):
x = [t + p for t, p in zip(W['wte'][token_id], W['wpe'][pos_id])] # 8.1
x = rmsnorm(x) # 8.5
# TODO: for li in range(N_LAYER):
# 1) the attention block (8.2 to 9.1):
# keep x_residual = x, then rmsnorm(x)
# q from attn_wq; append this token's key (attn_wk) to keys[li], its value (attn_wv) to values[li]
# x = attn_wo applied to multi_head(q, keys[li], values[li])
# add x_residual back, number by number (see the page; x + x_residual glues the lists)
# 2) the MLP block (9.2): keep x_residual = x, x = mlp(rmsnorm(x), li), add x_residual back
# TODO: lm_head turns the 16 numbers into 27 scores (6.5)
return x
# Use it: what does the trained model expect after "ann"?
keys, values = [[] for _ in range(N_LAYER)], [[] for _ in range(N_LAYER)]
for pos, tok in enumerate([BOS] + [UCHARS.index(c) for c in 'ann']):
logits = gpt(tok, pos, keys, values)
print('gpt() returned', len(logits), 'scores (should be 27)')
probs = softmax(logits)
top = sorted(range(len(probs)), key=lambda i: -probs[i])[:5]
print('after "ann":', [((UCHARS + '.')[i], round(probs[i], 3)) for i in top])
4Check yourself
5Unlocked in microgpt
Lines 92 and 93 (the architecture comments), line 108 (the function itself) and lines 114 and 115 (the layer loop). With these, the whole model, lines 92 to 144, is yours.
Progress is saved in this browser. Sign in to keep it across devices.