sweedworks

← all sources

checkattn.py

Gradient checks for the attention head

123 lines. This is the file the build actually runs, copied verbatim at build time.

  1. 1"""Check the attention head in the browser against the Python reference.
  2. 2 
  3. 3Same arrangement as checkmlp.py, and for the same honest reason: training is too
  4. 4deep in floating point for bit-identical output across engines. So each side
  5. 5checks its own analytic gradients against finite differences — the real risk in
  6. 6hand-written backprop through a softmax — and the two are compared on
  7. 7initialisation, forward pass and attention weights, which are shallow enough to
  8. 8agree exactly.
  9. 9"""
  10. 10 
  11. 11import json
  12. 12import os
  13. 13import sys
  14. 14 
  15. 15HERE = os.path.dirname(os.path.abspath(__file__))
  16. 16sys.path.insert(0, os.path.join(HERE, "pylib"))
  17. 17ROOT = os.path.abspath(os.path.join(HERE, ".."))
  18. 18 
  19. 19import quickjs # noqa: E402
  20. 20 
  21. 21import attn # noqa: E402
  22. 22import task # noqa: E402
  23. 23from corpora import BERRIES # noqa: E402
  24. 24 
  25. 25TOL = 1e-9
  26. 26 
  27. 27 
  28. 28def js_context():
  29. 29 ctx = quickjs.Context()
  30. 30 ctx.set_memory_limit(1 << 30)
  31. 31 ctx.set_max_stack_size(1 << 22)
  32. 32 ctx.eval("var globalThis = this;")
  33. 33 ctx.eval(open(os.path.join(ROOT, "attention", "attn.js"),
  34. 34 encoding="utf-8").read())
  35. 35 return ctx
  36. 36 
  37. 37 
  38. 38def main():
  39. 39 ctx = js_context()
  40. 40 failures = []
  41. 41 
  42. 42 copy_text, _ = task.make_corpus(120, seed=1)
  43. 43 cases = [("copy-task", copy_text, True), ("copy-task-plain", copy_text, False),
  44. 44 ("berries", BERRIES, True)]
  45. 45 
  46. 46 for name, text, shift in cases:
  47. 47 vocab = attn.build_vocab(text)
  48. 48 xs, ys = attn.make_examples(text, vocab)
  49. 49 model = attn.Model(vocab, seed=1, shift_values=shift)
  50. 50 
  51. 51 ctx.set("_t", text)
  52. 52 ctx.set("_shift", shift)
  53. 53 ctx.eval("var vocab = Attn.buildVocab(_t);"
  54. 54 "var ex = Attn.makeExamples(_t, vocab);"
  55. 55 "var m = new Attn.Model(vocab, 1, {shiftValues: _shift});")
  56. 56 
  57. 57 if json.loads(ctx.eval("JSON.stringify(vocab)")) != vocab:
  58. 58 failures.append(f"{name}: vocabulary differs")
  59. 59 continue
  60. 60 if json.loads(ctx.eval("JSON.stringify(ex.xs)")) != xs:
  61. 61 failures.append(f"{name}: examples differ")
  62. 62 continue
  63. 63 if int(ctx.eval("m.parameterCount()")) != model.parameter_count():
  64. 64 failures.append(f"{name}: parameter count differs")
  65. 65 
  66. 66 # Initialisation
  67. 67 worst_init = 0.0
  68. 68 for tensor in ("C", "P", "Wq", "Wk", "Wv", "Wo"):
  69. 69 js = json.loads(ctx.eval(f"JSON.stringify(m.{tensor})"))
  70. 70 py = getattr(model, tensor)
  71. 71 for a_row, b_row in zip(py, js):
  72. 72 for a, b in zip(a_row, b_row):
  73. 73 worst_init = max(worst_init, abs(a - b))
  74. 74 if worst_init > TOL:
  75. 75 failures.append(f"{name}: initial weights differ by {worst_init:.2e}")
  76. 76 
  77. 77 # Forward pass and, importantly, the attention weights themselves —
  78. 78 # they are what the page shows the reader.
  79. 79 worst_fwd = worst_attn = 0.0
  80. 80 for probe in xs[:20]:
  81. 81 py_probs, py_cache = model.forward(probe)
  82. 82 ctx.set("_c", json.dumps(probe))
  83. 83 ctx.eval("var out = m.forward(JSON.parse(_c));")
  84. 84 js_probs = json.loads(ctx.eval("JSON.stringify(out.probs)"))
  85. 85 js_w = json.loads(ctx.eval("JSON.stringify(out.w)"))
  86. 86 for a, b in zip(py_probs, js_probs):
  87. 87 worst_fwd = max(worst_fwd, abs(a - b))
  88. 88 for a, b in zip(py_cache["w"], js_w):
  89. 89 worst_attn = max(worst_attn, abs(a - b))
  90. 90 if worst_fwd > TOL:
  91. 91 failures.append(f"{name}: forward differs by {worst_fwd:.2e}")
  92. 92 if worst_attn > TOL:
  93. 93 failures.append(f"{name}: attention weights differ by {worst_attn:.2e}")
  94. 94 
  95. 95 # Each side checks its own derivatives.
  96. 96 py_worst, py_fail = attn.gradcheck(model, xs[:6], ys[:6])
  97. 97 ctx.eval("var gc = m.gradcheck(ex.xs.slice(0,6), ex.ys.slice(0,6), "
  98. 98 "{checks: 40});")
  99. 99 js_worst = float(ctx.eval("gc.worst"))
  100. 100 js_fail = int(ctx.eval("gc.failures.length"))
  101. 101 if py_fail:
  102. 102 failures.append(f"{name}: python gradients wrong ({len(py_fail)})")
  103. 103 if js_fail:
  104. 104 failures.append(f"{name}: browser gradients wrong ({js_fail})")
  105. 105 
  106. 106 print(f"PASS {name:<16} grad py {py_worst:.1e} / js {js_worst:.1e} "
  107. 107 f"init {worst_init:.1e} fwd {worst_fwd:.1e} "
  108. 108 f"attn {worst_attn:.1e}")
  109. 109 
  110. 110 print()
  111. 111 if failures:
  112. 112 print(f"{len(failures)} problem(s):")
  113. 113 for f in failures:
  114. 114 print(f" - {f}")
  115. 115 return 1
  116. 116 print("Browser attention head and Python reference agree; "
  117. 117 "both gradient checks pass.")
  118. 118 return 0
  119. 119 
  120. 120 
  121. 121if __name__ == "__main__":
  122. 122 sys.exit(main())