checkngram.py
Browser sampler vs the Python reference
158 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""Assert predict/ngram.js and .build/ngram.py are the same model. - 2
- 3
Three layers, because each can break independently: - 4
- 5
1. the PRNG stream — a seeded generator is worthless if the two languages - 6
disagree on the 32-bit arithmetic, and JS coercions vs Python masking is - 7
exactly where that goes wrong - 8
2. the distributions, before and after temperature / top-k / top-p - 9
3. whole generated sequences, which will diverge on the first mismatch of - 10
either of the above - 11
- 12
Probabilities are compared with a tolerance, because pow() may differ in the - 13
last bit between implementations. Chosen tokens are compared exactly: if a - 14
last-bit difference ever flips a choice, that is worth knowing about. - 15
""" - 16
- 17
import json - 18
import os - 19
import sys - 20
- 21
HERE = os.path.dirname(os.path.abspath(__file__)) - 22
sys.path.insert(0, os.path.join(HERE, "pylib")) - 23
ROOT = os.path.abspath(os.path.join(HERE, "..")) - 24
- 25
import quickjs # noqa: E402 - 26
- 27
import ngram # noqa: E402 - 28
from corpora import CORPORA, TALE # noqa: E402 - 29
- 30
TOL = 1e-12 - 31
- 32
SETTINGS = [ - 33
{"temp": 1.0, "top_k": 0, "top_p": 0.0}, - 34
{"temp": 0.0, "top_k": 0, "top_p": 0.0}, # greedy - 35
{"temp": 0.5, "top_k": 0, "top_p": 0.0}, - 36
{"temp": 2.5, "top_k": 0, "top_p": 0.0}, - 37
{"temp": 1.0, "top_k": 3, "top_p": 0.0}, - 38
{"temp": 1.0, "top_k": 0, "top_p": 0.5}, - 39
{"temp": 1.7, "top_k": 5, "top_p": 0.9}, - 40
] - 41
- 42
- 43
def js_context(): - 44
ctx = quickjs.Context() - 45
ctx.set_memory_limit(1 << 29) - 46
ctx.set_max_stack_size(1 << 22) - 47
ctx.eval("var globalThis = this;") - 48
ctx.eval(open(os.path.join(ROOT, "predict", "ngram.js"), encoding="utf-8").read()) - 49
return ctx - 50
- 51
- 52
def close(a, b): - 53
return abs(a - b) <= TOL - 54
- 55
- 56
def main(): - 57
ctx = js_context() - 58
failures = [] - 59
- 60
# 1. PRNG streams must be bit-identical. - 61
for seed in [0, 1, 7, 12345, 2 ** 31, 4294967295]: - 62
py = [ngram.Rng(seed).next() for _ in range(1)] - 63
r = ngram.Rng(seed) - 64
py = [r.next() for _ in range(2000)] - 65
ctx.set("_seed", seed) - 66
js = json.loads(ctx.eval( - 67
"(function(){var r=new NGram.Rng(_seed),o=[];" - 68
"for(var i=0;i<2000;i++)o.push(r.next());return JSON.stringify(o)})()")) - 69
bad = [i for i, (a, b) in enumerate(zip(py, js)) if a != b] - 70
if bad: - 71
failures.append(f"PRNG seed {seed}: differs at draw {bad[0]} " - 72
f"({py[bad[0]]} vs {js[bad[0]]})") - 73
if not failures: - 74
print("PASS PRNG — 6 seeds x 2000 draws, bit-identical") - 75
- 76
corpora = [("tale", TALE, 2)] + [(n, t, 2) for n, t, _ in CORPORA] - 77
- 78
for name, text, order in corpora: - 79
for o in (order, 1, 0): - 80
py_model = ngram.train(text, o) - 81
ctx.set("_t", text) - 82
ctx.set("_o", o) - 83
ctx.eval("var M = NGram.train(_t, _o);") - 84
- 85
if json.loads(ctx.eval("JSON.stringify(M.vocab)")) != py_model["vocab"]: - 86
failures.append(f"{name}/order{o}: vocabulary differs") - 87
continue - 88
- 89
# 2. distributions, raw and shaped - 90
probes = [py_model["tokens"][:i] for i in (0, 1, 2, 3)] - 91
for probe in probes: - 92
py_items, py_order = ngram.distribution(py_model, probe) - 93
ctx.set("_c", json.dumps(probe)) - 94
ctx.eval("var D = NGram.distribution(M, JSON.parse(_c));") - 95
js_items = json.loads(ctx.eval("JSON.stringify(D.items)")) - 96
js_order = int(ctx.eval("D.order")) - 97
- 98
if js_order != py_order: - 99
failures.append(f"{name}/order{o}: backoff order " - 100
f"{js_order} vs {py_order}") - 101
break - 102
if [t for t, _ in py_items] != [t for t, _ in js_items]: - 103
failures.append(f"{name}/order{o}: distribution token order") - 104
break - 105
if any(not close(a[1], b[1]) for a, b in zip(py_items, js_items)): - 106
failures.append(f"{name}/order{o}: probabilities differ") - 107
break - 108
- 109
for s in SETTINGS: - 110
py_shaped = ngram.shape(py_items, s["temp"], s["top_k"], - 111
s["top_p"]) - 112
ctx.set("_temp", s["temp"]) - 113
ctx.set("_k", s["top_k"]) - 114
ctx.set("_p", s["top_p"]) - 115
js_shaped = json.loads(ctx.eval( - 116
"JSON.stringify(NGram.shape(D.items, _temp, _k, _p))")) - 117
if [t for t, _ in py_shaped] != [t for t, _ in js_shaped]: - 118
failures.append(f"{name}/order{o} {s}: shaped tokens differ") - 119
break - 120
if any(not close(a[1], b[1]) - 121
for a, b in zip(py_shaped, js_shaped)): - 122
failures.append(f"{name}/order{o} {s}: shaped probs differ") - 123
break - 124
- 125
# 3. whole generated sequences - 126
for s in SETTINGS: - 127
for seed in (3, 7, 99): - 128
py_gen = ngram.generate(py_model, "", 40, s["temp"], - 129
s["top_k"], s["top_p"], seed) - 130
ctx.set("_temp", s["temp"]) - 131
ctx.set("_k", s["top_k"]) - 132
ctx.set("_p", s["top_p"]) - 133
ctx.set("_seed", seed) - 134
js_text = ctx.eval( - 135
"NGram.generate(M, '', 40, _temp, _k, _p, _seed).text") - 136
if js_text != py_gen["text"]: - 137
failures.append( - 138
f"{name}/order{o} {s} seed {seed}: generated text " - 139
f"diverges\n js {js_text[:60]!r}\n py " - 140
f"{py_gen['text'][:60]!r}") - 141
break - 142
- 143
print(f"PASS {name:<14} orders {order}/1/0, {len(SETTINGS)} settings, " - 144
f"3 seeds") - 145
- 146
print() - 147
if failures: - 148
print(f"{len(failures)} disagreement(s):") - 149
for f in failures[:8]: - 150
print(f" - {f}") - 151
return 1 - 152
print(f"Browser sampler and Python reference agree on {len(corpora)} corpora.") - 153
return 0 - 154
- 155
- 156
if __name__ == "__main__": - 157
sys.exit(main())