8.4 Causal masking & the KV cache
When you predict the next letter of a name, the letters after it don’t exist yet. A GPT must never peek at the future, not even during training. microgpt enforces that rule with a list that only ever grows.
1Watch
microgpt reads a name one token at a time. At each position it makes this token’s key and value and appends them to two lists, the KV cache. Then the token’s query attends over the cache.
Because the cache only contains tokens seen so far, position 3 can attend to positions 0 to 3 and nothing else. That’s causal masking for free: the triangle shape in every heatmap. A nice consequence is that adding letters to the end never changes what earlier positions computed. Your lab checks exactly that.
Why “masking”? Big GPTs read the whole name at once, for speed, so they hide the future with a mask: future scores are set to −∞, which softmax turns into 0%. microgpt reads one token at a time, so the future simply isn’t in the list yet. Same rule, no mask needed.
The same cache makes generating text fast: each new letter only computes its own key and value instead of redoing the whole name. (Your lab’s reader is given an rmsnorm, which you build in 8.5, so it sees exactly what the real model sees.)
2Explore
3Build
One blank in read(name, head): after computing this token’s query, append its key and value to the cache. Use the same pattern as q, with the key and value matrices and the same [s:e] slice, so each head keeps only its own 4 numbers. The lab checks head 1 and head 3 against the real model.
Below it you’re given read_cheating, which makes every key and value of the whole name first. Run the lab and compare: your reader’s early rows don’t change when “em” grows to “emma”, but the cheater’s do. The future leaked in. That is what the cache prevents.
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]
def rmsnorm(x): # given here; you build it in lesson 8.5
ms = sum(v * v for v in x) / len(x)
return [v * (ms + 1e-5) ** -0.5 for v in x]
def attend(q, keys, values): # lesson 8.3
# assert cond, message: carry on if cond is True, otherwise stop with message
assert keys, ("the cache is empty, so there is nothing to attend to yet: append this "
"token's key and value first (the TODO in read)")
scores = [sum(qi * ki for qi, ki in zip(q, k)) / math.sqrt(len(q)) for k in keys]
weights = softmax(scores)
return [sum(w * v[j] for w, v in zip(weights, values)) for j in range(len(values[0]))], weights
# Read a name one token at a time. Each new token adds its key and value to
# the cache, then attends over everything cached so far - never the future.
def read(name, head=0):
tokens = [BOS] + [UCHARS.index(c) for c in name] # .index(c) = where c sits in UCHARS: its token id
keys, values = [], [] # the KV cache, empty at the start
all_weights, outputs = [], []
s, e = head * 4, head * 4 + 4 # this head's 4 numbers
for pos, tok in enumerate(tokens): # enumerate gives (0, first), (1, second), ...: position and item
x = [t + p for t, p in zip(W['wte'][tok], W['wpe'][pos])]
x = rmsnorm(rmsnorm(x)) # exactly as microgpt does before attention (lines 112, 117)
q = linear(x, W['layer0.attn_wq'])[s:e]
# TODO: append this token's key and value to the cache (same pattern as q,
# with W['layer0.attn_wk'] and W['layer0.attn_wv'], and the same [s:e] slice)
out, weights = attend(q, keys, values)
all_weights.append(weights)
outputs.append(out)
return all_weights, outputs
# Given, for comparison: a reader that CHEATS. It makes the keys and values of
# the WHOLE name first, then lets every position attend over all of them,
# future letters included. This is exactly what the cache prevents.
def read_cheating(name, head=0):
tokens = [BOS] + [UCHARS.index(c) for c in name]
s, e = head * 4, head * 4 + 4
xs = [rmsnorm(rmsnorm([t + p for t, p in zip(W['wte'][tok], W['wpe'][pos])]))
for pos, tok in enumerate(tokens)]
keys = [linear(x, W['layer0.attn_wk'])[s:e] for x in xs]
values = [linear(x, W['layer0.attn_wv'])[s:e] for x in xs]
all_weights, outputs = [], []
for x in xs:
out, weights = attend(linear(x, W['layer0.attn_wq'])[s:e], keys, values)
all_weights.append(weights)
outputs.append(out)
return all_weights, outputs
# No peeking? Read "em", then "emma". The first 3 rows (., e, m) should not change.
def biggest_change(reader):
short, _ = reader('em')
longer, _ = reader('emma')
biggest = 0.0
for r1, r2 in zip(short, longer):
for a, b in zip(r1, r2):
biggest = max(biggest, abs(a - b)) # abs(x) drops the sign: abs(-0.3) is 0.3
return biggest
weights, _ = read('emma')
for pos, w in enumerate(weights):
print('.emma'[pos], [round(x, 2) for x in w])
print('biggest change to the first 3 rows when "em" grows to "emma":')
print(' read ', round(biggest_change(read), 4))
print(' read_cheating ', round(biggest_change(read_cheating), 4), ' <- the future leaked in')
4Check yourself
5Unlocked in microgpt
Line 161 creates a fresh, empty cache (keys, values = …) for each name. Lines 121 and 122, which you met in 8.2, append to it. The training loop then feeds tokens one position at a time, exactly like your read().
Progress is saved in this browser. Sign in to keep it across devices.