sweedworks

← all sources

checkngram.py

Browser sampler vs the Python reference

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

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