mlp.py
The neural language model behind /learn/, by hand
311 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""A small neural language model, written out by hand. - 2
- 3
The n-gram on /predict/ cannot answer a question it was never asked: given a - 4
context it has not seen, it backs off to a shorter one and eventually to noise. - 5
This model can, because it does not store contexts — it stores a vector for each - 6
character and learns a function of those vectors. Characters that behave alike - 7
end up near each other, and a context made of familiar parts is answerable even - 8
if that exact context never occurred. - 9
- 10
Architecture, deliberately the smallest thing that shows the point: - 11
- 12
context of K characters - 13
-> embedding lookup (V x D), concatenated (K*D) - 14
-> linear W1 (K*D x H) + b1, tanh (H) - 15
-> linear W2 (H x V) + b2 (V) - 16
-> softmax probabilities - 17
- 18
No autodiff and no matrix library: the forward pass and every gradient are - 19
written out so they can be read. verify_gradients() checks the analytic - 20
gradients against finite differences, which is the real correctness argument — - 21
float arithmetic will not reproduce bit-for-bit across languages, but a wrong - 22
derivative shows up immediately. - 23
- 24
Mirrored by learn/mlp.js. - 25
""" - 26
- 27
import math - 28
- 29
CONTEXT = 3 # characters of context - 30
EMBED = 8 # dimensions per character - 31
HIDDEN = 64 - 32
- 33
- 34
class Rng: - 35
"""mulberry32 again, so initialisation is reproducible in both languages.""" - 36
- 37
def __init__(self, seed): - 38
self.state = seed & 0xFFFFFFFF - 39
- 40
def next(self): - 41
self.state = (self.state + 0x6D2B79F5) & 0xFFFFFFFF - 42
t = self.state - 43
t = ((t ^ (t >> 15)) * (t | 1)) & 0xFFFFFFFF - 44
t = (t ^ (t + ((t ^ (t >> 7)) * (t | 61) & 0xFFFFFFFF))) & 0xFFFFFFFF - 45
return ((t ^ (t >> 14)) & 0xFFFFFFFF) / 4294967296.0 - 46
- 47
def normal(self, scale): - 48
"""Box-Muller, so weights start small and centred.""" - 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
"""Characters in first-appearance order, plus a boundary marker at index 0.""" - 56
vocab = ["\n"] - 57
seen = {"\n"} - 58
for ch in text: - 59
if ch not in seen: - 60
seen.add(ch) - 61
vocab.append(ch) - 62
return vocab - 63
- 64
- 65
def make_examples(text, vocab, context=CONTEXT): - 66
"""Every position becomes (context indices -> next index).""" - 67
index = {c: i for i, c in enumerate(vocab)} - 68
padded = "\n" * context + text - 69
xs, ys = [], [] - 70
for i in range(context, len(padded)): - 71
xs.append([index[padded[i - context + j]] for j in range(context)]) - 72
ys.append(index[padded[i]]) - 73
return xs, ys - 74
- 75
- 76
class Model: - 77
def __init__(self, vocab, seed=1, context=CONTEXT, embed=EMBED, hidden=HIDDEN): - 78
self.vocab = vocab - 79
self.V = len(vocab) - 80
self.K = context - 81
self.D = embed - 82
self.H = hidden - 83
rng = Rng(seed) - 84
self.C = [[rng.normal(1.0) for _ in range(self.D)] for _ in range(self.V)] - 85
fan_in = self.K * self.D - 86
s1 = 1.0 / math.sqrt(fan_in) - 87
self.W1 = [[rng.normal(s1) for _ in range(self.H)] for _ in range(fan_in)] - 88
self.b1 = [0.0] * self.H - 89
s2 = 1.0 / math.sqrt(self.H) - 90
self.W2 = [[rng.normal(s2) for _ in range(self.V)] for _ in range(self.H)] - 91
self.b2 = [0.0] * self.V - 92
- 93
# ---- forward ---- - 94
- 95
def forward(self, ctx): - 96
"""Returns (probabilities, cache) for one example.""" - 97
emb = [] - 98
for idx in ctx: - 99
emb.extend(self.C[idx]) - 100
- 101
h_pre = list(self.b1) - 102
for i, e in enumerate(emb): - 103
if e == 0.0: - 104
continue - 105
row = self.W1[i] - 106
for j in range(self.H): - 107
h_pre[j] += e * row[j] - 108
h = [math.tanh(v) for v in h_pre] - 109
- 110
logits = list(self.b2) - 111
for j in range(self.H): - 112
hj = h[j] - 113
row = self.W2[j] - 114
for k in range(self.V): - 115
logits[k] += hj * row[k] - 116
- 117
m = max(logits) - 118
exps = [math.exp(v - m) for v in logits] - 119
total = 0.0 - 120
for e in exps: - 121
total += e - 122
probs = [e / total for e in exps] - 123
return probs, {"emb": emb, "h": h, "ctx": ctx} - 124
- 125
def loss(self, xs, ys): - 126
total = 0.0 - 127
for ctx, y in zip(xs, ys): - 128
probs, _ = self.forward(ctx) - 129
total += -math.log(max(probs[y], 1e-12)) - 130
return total / len(xs) - 131
- 132
# ---- backward ---- - 133
- 134
def zero_grads(self): - 135
return { - 136
"C": [[0.0] * self.D for _ in range(self.V)], - 137
"W1": [[0.0] * self.H for _ in range(self.K * self.D)], - 138
"b1": [0.0] * self.H, - 139
"W2": [[0.0] * self.V for _ in range(self.H)], - 140
"b2": [0.0] * self.V, - 141
} - 142
- 143
def backward(self, xs, ys, grads): - 144
"""Accumulate gradients of mean cross-entropy over the batch.""" - 145
n = len(xs) - 146
total_loss = 0.0 - 147
for ctx, y in zip(xs, ys): - 148
probs, cache = self.forward(ctx) - 149
total_loss += -math.log(max(probs[y], 1e-12)) - 150
- 151
# dL/dlogits for softmax + cross-entropy is (p - onehot)/n - 152
dlogits = [p / n for p in probs] - 153
dlogits[y] -= 1.0 / n - 154
- 155
h = cache["h"] - 156
emb = cache["emb"] - 157
- 158
dh = [0.0] * self.H - 159
for j in range(self.H): - 160
row = self.W2[j] - 161
grow = grads["W2"][j] - 162
hj = h[j] - 163
acc = 0.0 - 164
for k in range(self.V): - 165
dk = dlogits[k] - 166
grow[k] += hj * dk - 167
acc += row[k] * dk - 168
dh[j] = acc - 169
for k in range(self.V): - 170
grads["b2"][k] += dlogits[k] - 171
- 172
# through tanh - 173
dh_pre = [dh[j] * (1.0 - h[j] * h[j]) for j in range(self.H)] - 174
- 175
demb = [0.0] * (self.K * self.D) - 176
for i in range(self.K * self.D): - 177
row = self.W1[i] - 178
grow = grads["W1"][i] - 179
e = emb[i] - 180
acc = 0.0 - 181
for j in range(self.H): - 182
dj = dh_pre[j] - 183
grow[j] += e * dj - 184
acc += row[j] * dj - 185
demb[i] = acc - 186
for j in range(self.H): - 187
grads["b1"][j] += dh_pre[j] - 188
- 189
for slot, idx in enumerate(cache["ctx"]): - 190
base = slot * self.D - 191
grow = grads["C"][idx] - 192
for d in range(self.D): - 193
grow[d] += demb[base + d] - 194
- 195
return total_loss / n - 196
- 197
def step(self, grads, lr): - 198
for i in range(self.V): - 199
row, g = self.C[i], grads["C"][i] - 200
for d in range(self.D): - 201
row[d] -= lr * g[d] - 202
for i in range(self.K * self.D): - 203
row, g = self.W1[i], grads["W1"][i] - 204
for j in range(self.H): - 205
row[j] -= lr * g[j] - 206
for j in range(self.H): - 207
self.b1[j] -= lr * grads["b1"][j] - 208
row, g = self.W2[j], grads["W2"][j] - 209
for k in range(self.V): - 210
row[k] -= lr * g[k] - 211
for k in range(self.V): - 212
self.b2[k] -= lr * grads["b2"][k] - 213
- 214
def train(self, xs, ys, steps, lr, batch=32, seed=7, on_step=None): - 215
rng = Rng(seed) - 216
history = [] - 217
for s in range(steps): - 218
idx = [int(rng.next() * len(xs)) % len(xs) for _ in range(batch)] - 219
bx = [xs[i] for i in idx] - 220
by = [ys[i] for i in idx] - 221
grads = self.zero_grads() - 222
loss = self.backward(bx, by, grads) - 223
self.step(grads, lr) - 224
history.append(loss) - 225
if on_step: - 226
on_step(s, loss) - 227
return history - 228
- 229
- 230
def verify_gradients(model, xs, ys, eps=1e-5, tol=1e-4, checks=40, seed=3): - 231
"""Analytic gradients vs finite differences. - 232
- 233
This is the correctness argument for the whole page. Float arithmetic will - 234
not agree bit-for-bit between Python and JavaScript, so "both implementations - 235
match exactly" is not available here — but a wrong derivative fails this - 236
immediately, in either language. - 237
""" - 238
grads = model.zero_grads() - 239
model.backward(xs, ys, grads) - 240
rng = Rng(seed) - 241
- 242
params = [ - 243
("C", model.C, grads["C"], True), - 244
("W1", model.W1, grads["W1"], True), - 245
("W2", model.W2, grads["W2"], True), - 246
("b1", model.b1, grads["b1"], False), - 247
("b2", model.b2, grads["b2"], False), - 248
] - 249
- 250
worst = 0.0 - 251
failures = [] - 252
for _ in range(checks): - 253
name, tensor, grad, two_d = params[int(rng.next() * len(params))] - 254
if two_d: - 255
i = int(rng.next() * len(tensor)) - 256
j = int(rng.next() * len(tensor[i])) - 257
original = tensor[i][j] - 258
analytic = grad[i][j] - 259
tensor[i][j] = original + eps - 260
plus = model.loss(xs, ys) - 261
tensor[i][j] = original - eps - 262
minus = model.loss(xs, ys) - 263
tensor[i][j] = original - 264
else: - 265
i = int(rng.next() * len(tensor)) - 266
original = tensor[i] - 267
analytic = grad[i] - 268
tensor[i] = original + eps - 269
plus = model.loss(xs, ys) - 270
tensor[i] = original - eps - 271
minus = model.loss(xs, ys) - 272
tensor[i] = original - 273
- 274
numeric = (plus - minus) / (2 * eps) - 275
scale = max(abs(analytic), abs(numeric), 1e-8) - 276
rel = abs(analytic - numeric) / scale - 277
worst = max(worst, rel) - 278
if rel > tol: - 279
failures.append((name, analytic, numeric, rel)) - 280
- 281
return worst, failures - 282
- 283
- 284
def generate(model, prompt, n, temp=1.0, seed=11): - 285
"""Sample n characters. Shares the sampling arithmetic from /predict/.""" - 286
rng = Rng(seed) - 287
index = {c: i for i, c in enumerate(model.vocab)} - 288
ctx = [index.get(c, 0) for c in ("\n" * model.K + prompt)[-model.K:]] - 289
out = [] - 290
for _ in range(n): - 291
probs, _ = model.forward(ctx) - 292
if temp <= 0: - 293
pick = max(range(len(probs)), key=lambda i: probs[i]) - 294
else: - 295
scaled = [p ** (1.0 / temp) for p in probs] - 296
total = 0.0 - 297
for v in scaled: - 298
total += v - 299
scaled = [v / total for v in scaled] - 300
r = rng.next() - 301
acc = 0.0 - 302
pick = len(scaled) - 1 - 303
for i, v in enumerate(scaled): - 304
acc += v - 305
if r < acc: - 306
pick = i - 307
break - 308
out.append(model.vocab[pick]) - 309
ctx = ctx[1:] + [pick] - 310
return "".join(out)