sweedworks

← all sources

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