sweedworks

← all sources

checkmlp.py

Gradient checks, and browser vs Python agreement

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

  1. 1"""Check the neural model in the browser against the Python reference.
  2. 2 
  3. 3Bit-identical output is not available here and pretending otherwise would be
  4. 4dishonest: training runs thousands of floating-point operations deep, and
  5. 5tanh, exp and log differ in their last bits between engines. So this checks
  6. 6three things instead, which together are a stronger claim than "the numbers
  7. 7matched once":
  8. 8 
  9. 9 1. Each implementation's analytic gradients match its own finite differences.
  10. 10 A wrong derivative — the actual risk in hand-written backprop — fails here
  11. 11 immediately, in whichever language it was written in.
  12. 12 2. The two agree to a tolerance on initialisation and the forward pass, so
  13. 13 they are the same model and not merely two models that both train.
  14. 14 3. Training reduces loss in both, from the same start, by a similar amount.
  15. 15"""
  16. 16 
  17. 17import json
  18. 18import math
  19. 19import os
  20. 20import sys
  21. 21 
  22. 22HERE = os.path.dirname(os.path.abspath(__file__))
  23. 23sys.path.insert(0, os.path.join(HERE, "pylib"))
  24. 24ROOT = os.path.abspath(os.path.join(HERE, ".."))
  25. 25 
  26. 26import quickjs # noqa: E402
  27. 27 
  28. 28import mlp # noqa: E402
  29. 29from corpora import BERRIES, CODE, JAPANESE # noqa: E402
  30. 30 
  31. 31CORPORA = [("berries", BERRIES), ("code", CODE), ("japanese", JAPANESE)]
  32. 32 
  33. 33INIT_TOL = 1e-9 # same PRNG, same order: only transcendentals can differ
  34. 34FORWARD_TOL = 1e-9
  35. 35GRAD_TOL = 1e-4 # finite differences are not exact; this is the usual bar
  36. 36 
  37. 37 
  38. 38def js_context():
  39. 39 ctx = quickjs.Context()
  40. 40 ctx.set_memory_limit(1 << 30)
  41. 41 ctx.set_max_stack_size(1 << 22)
  42. 42 ctx.eval("var globalThis = this;")
  43. 43 ctx.eval(open(os.path.join(ROOT, "learn", "mlp.js"), encoding="utf-8").read())
  44. 44 return ctx
  45. 45 
  46. 46 
  47. 47def main():
  48. 48 ctx = js_context()
  49. 49 failures = []
  50. 50 
  51. 51 for name, text in CORPORA:
  52. 52 vocab = mlp.build_vocab(text)
  53. 53 xs, ys = mlp.make_examples(text, vocab)
  54. 54 model = mlp.Model(vocab, seed=1)
  55. 55 
  56. 56 ctx.set("_t", text)
  57. 57 ctx.eval("var vocab = MLP.buildVocab(_t);"
  58. 58 "var ex = MLP.makeExamples(_t, vocab);"
  59. 59 "var m = new MLP.Model(vocab, 1);")
  60. 60 
  61. 61 js_vocab = json.loads(ctx.eval("JSON.stringify(vocab)"))
  62. 62 if js_vocab != vocab:
  63. 63 failures.append(f"{name}: vocabulary differs "
  64. 64 f"({len(js_vocab)} vs {len(vocab)})")
  65. 65 continue
  66. 66 js_xs = json.loads(ctx.eval("JSON.stringify(ex.xs)"))
  67. 67 js_ys = json.loads(ctx.eval("JSON.stringify(ex.ys)"))
  68. 68 if js_xs != xs or js_ys != ys:
  69. 69 failures.append(f"{name}: training examples differ")
  70. 70 continue
  71. 71 
  72. 72 # 2a. same initial weights
  73. 73 js_C = json.loads(ctx.eval("JSON.stringify(m.C)"))
  74. 74 worst_init = 0.0
  75. 75 for a_row, b_row in zip(model.C, js_C):
  76. 76 for a, b in zip(a_row, b_row):
  77. 77 worst_init = max(worst_init, abs(a - b))
  78. 78 js_W2 = json.loads(ctx.eval("JSON.stringify(m.W2)"))
  79. 79 for a_row, b_row in zip(model.W2, js_W2):
  80. 80 for a, b in zip(a_row, b_row):
  81. 81 worst_init = max(worst_init, abs(a - b))
  82. 82 if worst_init > INIT_TOL:
  83. 83 failures.append(f"{name}: initial weights differ by {worst_init:.2e}")
  84. 84 
  85. 85 # 2b. same forward pass
  86. 86 worst_fwd = 0.0
  87. 87 for probe in xs[:25]:
  88. 88 py_probs, _ = model.forward(probe)
  89. 89 ctx.set("_c", json.dumps(probe))
  90. 90 js_probs = json.loads(ctx.eval(
  91. 91 "JSON.stringify(m.forward(JSON.parse(_c)).probs)"))
  92. 92 for a, b in zip(py_probs, js_probs):
  93. 93 worst_fwd = max(worst_fwd, abs(a - b))
  94. 94 if worst_fwd > FORWARD_TOL:
  95. 95 failures.append(f"{name}: forward pass differs by {worst_fwd:.2e}")
  96. 96 
  97. 97 # 1. each side checks its own gradients
  98. 98 py_worst, py_fail = mlp.verify_gradients(model, xs[:10], ys[:10])
  99. 99 ctx.eval("var sub = {xs: ex.xs.slice(0,10), ys: ex.ys.slice(0,10)};"
  100. 100 "var gc = m.gradcheck(sub.xs, sub.ys, {checks: 40});")
  101. 101 js_worst = float(ctx.eval("gc.worst"))
  102. 102 js_fail = int(ctx.eval("gc.failures.length"))
  103. 103 if py_fail:
  104. 104 failures.append(f"{name}: python gradients wrong ({len(py_fail)})")
  105. 105 if js_fail:
  106. 106 failures.append(f"{name}: browser gradients wrong ({js_fail})")
  107. 107 
  108. 108 # 3. both actually learn
  109. 109 py_before = model.loss(xs[:120], ys[:120])
  110. 110 model.train(xs, ys, 120, 0.5, seed=7)
  111. 111 py_after = model.loss(xs[:120], ys[:120])
  112. 112 
  113. 113 ctx.eval("var before = m.loss(ex.xs.slice(0,120), ex.ys.slice(0,120));"
  114. 114 "var rng = new MLP.Rng(7);"
  115. 115 "m.train(ex.xs, ex.ys, 120, 0.5, 32, rng);"
  116. 116 "var after = m.loss(ex.xs.slice(0,120), ex.ys.slice(0,120));")
  117. 117 js_before = float(ctx.eval("before"))
  118. 118 js_after = float(ctx.eval("after"))
  119. 119 
  120. 120 if not (py_after < py_before and js_after < js_before):
  121. 121 failures.append(f"{name}: loss did not fall "
  122. 122 f"(py {py_before:.3f}->{py_after:.3f}, "
  123. 123 f"js {js_before:.3f}->{js_after:.3f})")
  124. 124 
  125. 125 print(f"PASS {name:<10} grad max rel err py {py_worst:.1e} / "
  126. 126 f"js {js_worst:.1e} init {worst_init:.1e} "
  127. 127 f"fwd {worst_fwd:.1e} loss {py_before:.2f}->{py_after:.2f}")
  128. 128 
  129. 129 print()
  130. 130 if failures:
  131. 131 print(f"{len(failures)} problem(s):")
  132. 132 for f in failures:
  133. 133 print(f" - {f}")
  134. 134 return 1
  135. 135 print("Browser model and Python reference agree; both gradient checks pass.")
  136. 136 return 0
  137. 137 
  138. 138 
  139. 139if __name__ == "__main__":
  140. 140 sys.exit(main())