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