sweedworks

← all sources

checkbpe.py

Browser BPE trainer vs the Python reference

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

  1. 1"""Assert the browser's BPE trainer and the Python reference agree exactly.
  2. 2 
  3. 3The page claims the algorithm you run in your browser is the one that produced
  4. 4the static figures. That claim is only worth making if it is checked, so this
  5. 5compares merge lists, step counts and encodings across a range of corpora —
  6. 6including the ones designed to break naive implementations: astral-plane
  7. 7characters, combining marks, and text with no repetition at all.
  8. 8"""
  9. 9 
  10. 10import json
  11. 11import os
  12. 12import sys
  13. 13 
  14. 14HERE = os.path.dirname(os.path.abspath(__file__))
  15. 15sys.path.insert(0, os.path.join(HERE, "pylib"))
  16. 16ROOT = os.path.abspath(os.path.join(HERE, ".."))
  17. 17 
  18. 18import quickjs # noqa: E402
  19. 19 
  20. 20import bpe # noqa: E402
  21. 21from corpora import CORPORA # noqa: E402
  22. 22 
  23. 23PROBES = [" strawberry", "strawberry", " raspberry", "strawberries", " kiwi",
  24. 24 "", " ", "\n", "🍓🍓", "café", "the the the"]
  25. 25 
  26. 26 
  27. 27def js_context():
  28. 28 ctx = quickjs.Context()
  29. 29 ctx.set_memory_limit(1 << 29)
  30. 30 ctx.set_max_stack_size(1 << 22)
  31. 31 ctx.eval("var globalThis = this;")
  32. 32 ctx.eval(open(os.path.join(ROOT, "vocabulary", "bpe.js"), encoding="utf-8").read())
  33. 33 return ctx
  34. 34 
  35. 35 
  36. 36def main():
  37. 37 ctx = js_context()
  38. 38 failures = []
  39. 39 
  40. 40 for name, text, n_merges in CORPORA:
  41. 41 want = bpe.train(text, n_merges)
  42. 42 ctx.set("_t", text)
  43. 43 ctx.set("_n", n_merges)
  44. 44 got = json.loads(ctx.eval("JSON.stringify(BPE.train(_t, _n))"))
  45. 45 
  46. 46 w_merges = [list(m) for m in want["merges"]]
  47. 47 if got["merges"] != w_merges:
  48. 48 # Report the first divergence rather than dumping both lists.
  49. 49 for i, (a, b) in enumerate(zip(got["merges"], w_merges)):
  50. 50 if a != b:
  51. 51 failures.append(f"{name}: merge {i + 1} js={a} py={b}")
  52. 52 break
  53. 53 else:
  54. 54 failures.append(f"{name}: merge count js={len(got['merges'])} "
  55. 55 f"py={len(w_merges)}")
  56. 56 continue
  57. 57 
  58. 58 if got["vocabSize"] != want["vocab_size"]:
  59. 59 failures.append(f"{name}: vocab js={got['vocabSize']} "
  60. 60 f"py={want['vocab_size']}")
  61. 61 
  62. 62 if [s["count"] for s in got["steps"]] != [s["count"] for s in want["steps"]]:
  63. 63 failures.append(f"{name}: step counts differ")
  64. 64 
  65. 65 for probe in PROBES:
  66. 66 py = bpe.encode(probe, want["merges"])
  67. 67 ctx.set("_p", probe)
  68. 68 js = json.loads(ctx.eval(
  69. 69 "JSON.stringify(BPE.encode(_p, BPE.train(_t, _n).merges))"))
  70. 70 if js != py:
  71. 71 failures.append(f"{name}: encode({probe!r}) js={js} py={py}")
  72. 72 break
  73. 73 
  74. 74 print(f"PASS {name:<22} {len(w_merges):>3} merges, "
  75. 75 f"vocab {want['vocab_size']:>3}, {len(PROBES)} probes")
  76. 76 
  77. 77 # Pre-tokenization must agree too — it is the rule that welds on spaces.
  78. 78 for name, text, _ in CORPORA:
  79. 79 ctx.set("_t", text)
  80. 80 js = json.loads(ctx.eval("JSON.stringify(BPE.pretokenize(_t))"))
  81. 81 py = bpe.pretokenize(text)
  82. 82 if js != py:
  83. 83 failures.append(f"{name}: pretokenize differs "
  84. 84 f"({len(js)} vs {len(py)} chunks)")
  85. 85 
  86. 86 print()
  87. 87 if failures:
  88. 88 print(f"{len(failures)} disagreement(s) between trainer.js and bpe.py:")
  89. 89 for f in failures[:10]:
  90. 90 print(f" - {f}")
  91. 91 return 1
  92. 92 print(f"Browser trainer and Python reference agree exactly on "
  93. 93 f"{len(CORPORA)} corpora.")
  94. 94 return 0
  95. 95 
  96. 96 
  97. 97if __name__ == "__main__":
  98. 98 sys.exit(main())