attn.py
One attention head, and the induction-head result
349 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""One attention head, written out by hand. - 2
- 3
The model on /learn/ flattens a fixed window into one vector, so every position - 4
is wired to the output separately and the only way to use a character is to have - 5
learned a rule for that character *at that offset*. Attention replaces the - 6
flattening with a lookup by content: build a query from the current position, - 7
compare it against a key at every earlier position, and take a weighted average - 8
of the values there. Which position matters is decided at run time from the - 9
content, not baked into the weights. - 10
- 11
x_i = C[token_i] + P[i] embedding plus position - 12
q = x_last @ Wq one query, from where we are now - 13
k_i = x_i @ Wk a key at every position - 14
v_i = x_i @ Wv a value at every position - 15
score_i = (q . k_i) / sqrt(A) - 16
w = softmax(score) how much to look at each position - 17
context = sum_i w_i v_i - 18
logits = context @ Wo + b - 19
- 20
This is a single head at a single position — the smallest thing that shows the - 21
mechanism. A transformer does this at every position at once, several times in - 22
parallel, stacked in layers, with a feed-forward network between. - 23
- 24
No autodiff: every derivative is written out, and gradcheck() compares them - 25
against finite differences. Mirrored by attention/attn.js. - 26
""" - 27
- 28
import math - 29
- 30
WINDOW = 16 - 31
EMBED = 16 - 32
ATTN = 16 - 33
- 34
- 35
class Rng: - 36
"""mulberry32, same as everywhere else here.""" - 37
- 38
def __init__(self, seed): - 39
self.state = seed & 0xFFFFFFFF - 40
- 41
def next(self): - 42
self.state = (self.state + 0x6D2B79F5) & 0xFFFFFFFF - 43
t = self.state - 44
t = ((t ^ (t >> 15)) * (t | 1)) & 0xFFFFFFFF - 45
t = (t ^ (t + ((t ^ (t >> 7)) * (t | 61) & 0xFFFFFFFF))) & 0xFFFFFFFF - 46
return ((t ^ (t >> 14)) & 0xFFFFFFFF) / 4294967296.0 - 47
- 48
def normal(self, scale): - 49
u1 = max(self.next(), 1e-12) - 50
u2 = self.next() - 51
return scale * math.sqrt(-2 * math.log(u1)) * math.cos(2 * math.pi * u2) - 52
- 53
- 54
def build_vocab(text): - 55
vocab, seen = ["\n"], {"\n"} - 56
for ch in text: - 57
if ch not in seen: - 58
seen.add(ch) - 59
vocab.append(ch) - 60
return vocab - 61
- 62
- 63
def make_examples(text, vocab, window=WINDOW): - 64
index = {c: i for i, c in enumerate(vocab)} - 65
padded = "\n" * window + text - 66
xs, ys = [], [] - 67
for i in range(window, len(padded)): - 68
xs.append([index[padded[i - window + j]] for j in range(window)]) - 69
ys.append(index[padded[i]]) - 70
return xs, ys - 71
- 72
- 73
class Model: - 74
def __init__(self, vocab, seed=1, window=WINDOW, embed=EMBED, attn=ATTN, - 75
shift_values=False): - 76
self.vocab = vocab - 77
self.V = len(vocab) - 78
self.L = window - 79
self.D = embed - 80
self.A = attn - 81
rng = Rng(seed) - 82
s = 1.0 / math.sqrt(self.D) - 83
self.C = [[rng.normal(1.0) for _ in range(self.D)] for _ in range(self.V)] - 84
self.P = [[rng.normal(0.3) for _ in range(self.D)] for _ in range(self.L)] - 85
self.Wq = [[rng.normal(s) for _ in range(self.A)] for _ in range(self.D)] - 86
self.Wk = [[rng.normal(s) for _ in range(self.A)] for _ in range(self.D)] - 87
self.Wv = [[rng.normal(s) for _ in range(self.A)] for _ in range(self.D)] - 88
so = 1.0 / math.sqrt(self.A) - 89
self.Wo = [[rng.normal(so) for _ in range(self.V)] for _ in range(self.A)] - 90
self.b = [0.0] * self.V - 91
# When true, position i's value carries the NEXT character rather - 92
# than its own. A real transformer arranges this with an earlier - 93
# layer; supplying it lets one head show the matching half alone. - 94
self.shift_values = shift_values - 95
- 96
def parameter_count(self): - 97
return (self.V * self.D + self.L * self.D + 3 * self.D * self.A - 98
+ self.A * self.V + self.V) - 99
- 100
# ---- forward ---- - 101
- 102
def forward(self, ctx): - 103
L, D, A = self.L, self.D, self.A - 104
x = [] - 105
for i in range(L): - 106
emb = self.C[ctx[i]] - 107
pos = self.P[i] - 108
x.append([emb[d] + pos[d] for d in range(D)]) - 109
- 110
def project(vec, W): - 111
out = [0.0] * A - 112
for d in range(D): - 113
val = vec[d] - 114
if val == 0.0: - 115
continue - 116
row = W[d] - 117
for a in range(A): - 118
out[a] += val * row[a] - 119
return out - 120
- 121
q = project(x[L - 1], self.Wq) - 122
k = [project(x[i], self.Wk) for i in range(L)] - 123
if self.shift_values: - 124
zero = [0.0] * D - 125
xv = [x[i + 1] if i + 1 < L else zero for i in range(L)] - 126
else: - 127
xv = x - 128
v = [project(xv[i], self.Wv) for i in range(L)] - 129
- 130
inv = 1.0 / math.sqrt(A) - 131
scores = [] - 132
for i in range(L): - 133
dot = 0.0 - 134
ki = k[i] - 135
for a in range(A): - 136
dot += q[a] * ki[a] - 137
scores.append(dot * inv) - 138
- 139
m = max(scores) - 140
exps = [math.exp(s - m) for s in scores] - 141
total = 0.0 - 142
for e in exps: - 143
total += e - 144
w = [e / total for e in exps] - 145
- 146
context = [0.0] * A - 147
for i in range(L): - 148
wi = w[i] - 149
vi = v[i] - 150
for a in range(A): - 151
context[a] += wi * vi[a] - 152
- 153
logits = list(self.b) - 154
for a in range(A): - 155
ca = context[a] - 156
row = self.Wo[a] - 157
for c in range(self.V): - 158
logits[c] += ca * row[c] - 159
- 160
mm = max(logits) - 161
le = [math.exp(z - mm) for z in logits] - 162
lt = 0.0 - 163
for e in le: - 164
lt += e - 165
probs = [e / lt for e in le] - 166
- 167
return probs, {"x": x, "xv": xv, "q": q, "k": k, "v": v, "w": w, - 168
"context": context, "ctx": ctx} - 169
- 170
def loss(self, xs, ys): - 171
total = 0.0 - 172
for ctx, y in zip(xs, ys): - 173
probs, _ = self.forward(ctx) - 174
total += -math.log(max(probs[y], 1e-12)) - 175
return total / len(xs) - 176
- 177
def attention(self, ctx): - 178
"""The weights themselves — what the model looked at.""" - 179
_, cache = self.forward(ctx) - 180
return cache["w"] - 181
- 182
# ---- backward ---- - 183
- 184
def zero_grads(self): - 185
return { - 186
"C": [[0.0] * self.D for _ in range(self.V)], - 187
"P": [[0.0] * self.D for _ in range(self.L)], - 188
"Wq": [[0.0] * self.A for _ in range(self.D)], - 189
"Wk": [[0.0] * self.A for _ in range(self.D)], - 190
"Wv": [[0.0] * self.A for _ in range(self.D)], - 191
"Wo": [[0.0] * self.V for _ in range(self.A)], - 192
"b": [0.0] * self.V, - 193
} - 194
- 195
def backward(self, xs, ys, grads): - 196
L, D, A, V = self.L, self.D, self.A, self.V - 197
inv = 1.0 / math.sqrt(A) - 198
n = len(xs) - 199
total_loss = 0.0 - 200
- 201
for ctx, y in zip(xs, ys): - 202
probs, c = self.forward(ctx) - 203
total_loss += -math.log(max(probs[y], 1e-12)) - 204
x, xv, q, k, v, w, context = (c["x"], c["xv"], c["q"], c["k"], - 205
c["v"], c["w"], c["context"]) - 206
- 207
dlogits = [p / n for p in probs] - 208
dlogits[y] -= 1.0 / n - 209
- 210
dcontext = [0.0] * A - 211
for a in range(A): - 212
row = self.Wo[a] - 213
grow = grads["Wo"][a] - 214
ca = context[a] - 215
acc = 0.0 - 216
for cidx in range(V): - 217
dc = dlogits[cidx] - 218
grow[cidx] += ca * dc - 219
acc += row[cidx] * dc - 220
dcontext[a] = acc - 221
for cidx in range(V): - 222
grads["b"][cidx] += dlogits[cidx] - 223
- 224
# context = sum_i w_i v_i - 225
dw = [0.0] * L - 226
dv = [[0.0] * A for _ in range(L)] - 227
for i in range(L): - 228
vi = v[i] - 229
wi = w[i] - 230
acc = 0.0 - 231
dvi = dv[i] - 232
for a in range(A): - 233
acc += dcontext[a] * vi[a] - 234
dvi[a] = wi * dcontext[a] - 235
dw[i] = acc - 236
- 237
# through the softmax over scores - 238
dot = 0.0 - 239
for i in range(L): - 240
dot += w[i] * dw[i] - 241
dscore = [w[i] * (dw[i] - dot) for i in range(L)] - 242
- 243
# score_i = (q . k_i) * inv - 244
dq = [0.0] * A - 245
dk = [[0.0] * A for _ in range(L)] - 246
for i in range(L): - 247
s = dscore[i] * inv - 248
ki = k[i] - 249
dki = dk[i] - 250
for a in range(A): - 251
dq[a] += s * ki[a] - 252
dki[a] = s * q[a] - 253
- 254
dx = [[0.0] * D for _ in range(L)] - 255
- 256
def backprop_projection(vec, W, gW, dout, dvec): - 257
for d in range(D): - 258
val = vec[d] - 259
row = W[d] - 260
grow = gW[d] - 261
acc = 0.0 - 262
for a in range(A): - 263
da = dout[a] - 264
grow[a] += val * da - 265
acc += row[a] * da - 266
dvec[d] += acc - 267
- 268
backprop_projection(x[L - 1], self.Wq, grads["Wq"], dq, dx[L - 1]) - 269
for i in range(L): - 270
backprop_projection(x[i], self.Wk, grads["Wk"], dk[i], dx[i]) - 271
# values may read the next position, so the gradient goes there - 272
target = i + 1 if self.shift_values else i - 273
if target < L: - 274
backprop_projection(xv[i], self.Wv, grads["Wv"], dv[i], - 275
dx[target]) - 276
- 277
for i in range(L): - 278
gc = grads["C"][ctx[i]] - 279
gp = grads["P"][i] - 280
dxi = dx[i] - 281
for d in range(D): - 282
gc[d] += dxi[d] - 283
gp[d] += dxi[d] - 284
- 285
return total_loss / n - 286
- 287
def step(self, grads, lr): - 288
for name in ("C", "P", "Wq", "Wk", "Wv", "Wo"): - 289
tensor = getattr(self, name) - 290
grad = grads[name] - 291
for i in range(len(tensor)): - 292
row, g = tensor[i], grad[i] - 293
for j in range(len(row)): - 294
row[j] -= lr * g[j] - 295
for k in range(self.V): - 296
self.b[k] -= lr * grads["b"][k] - 297
- 298
def train(self, xs, ys, steps, lr, batch=32, seed=7): - 299
rng = Rng(seed) - 300
history = [] - 301
for _ in range(steps): - 302
idx = [int(rng.next() * len(xs)) % len(xs) for _ in range(batch)] - 303
bx = [xs[i] for i in idx] - 304
by = [ys[i] for i in idx] - 305
grads = self.zero_grads() - 306
history.append(self.backward(bx, by, grads)) - 307
self.step(grads, lr) - 308
return history - 309
- 310
- 311
def gradcheck(model, xs, ys, eps=1e-5, tol=1e-4, checks=40, seed=3): - 312
grads = model.zero_grads() - 313
model.backward(xs, ys, grads) - 314
rng = Rng(seed) - 315
names = ["C", "P", "Wq", "Wk", "Wv", "Wo"] - 316
worst, failures = 0.0, [] - 317
- 318
for _ in range(checks): - 319
if rng.next() < 0.12: - 320
i = int(rng.next() * model.V) - 321
original, analytic = model.b[i], grads["b"][i] - 322
model.b[i] = original + eps - 323
plus = model.loss(xs, ys) - 324
model.b[i] = original - eps - 325
minus = model.loss(xs, ys) - 326
model.b[i] = original - 327
label = "b" - 328
else: - 329
name = names[int(rng.next() * len(names))] - 330
tensor, grad = getattr(model, name), grads[name] - 331
i = int(rng.next() * len(tensor)) - 332
j = int(rng.next() * len(tensor[i])) - 333
original, analytic = tensor[i][j], grad[i][j] - 334
tensor[i][j] = original + eps - 335
plus = model.loss(xs, ys) - 336
tensor[i][j] = original - eps - 337
minus = model.loss(xs, ys) - 338
tensor[i][j] = original - 339
label = name - 340
- 341
numeric = (plus - minus) / (2 * eps) - 342
scale = max(abs(analytic), abs(numeric), 1e-8) - 343
rel = abs(analytic - numeric) / scale - 344
worst = max(worst, rel) - 345
if rel > tol: - 346
failures.append((label, analytic, numeric, rel)) - 347
- 348
return worst, failures