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