sweedworks

← all sources

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