sweedworks

← all sources

bpe.py

Byte pair encoding — the reference for /vocabulary/

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

  1. 1"""Byte pair encoding, trained from scratch — the reference implementation.
  2. 2 
  3. 3This is the teaching version of the algorithm that produced the vocabularies on
  4. 4/tokens/. It mirrors trainer.js exactly, symbol for symbol, and checkbpe.py
  5. 5fails the build if the two ever diverge.
  6. 6 
  7. 7Two honest simplifications versus a production tokenizer:
  8. 8 
  9. 9 * It merges characters, not UTF-8 bytes. Real tokenizers start from the 256
  10. 10 possible bytes so that any input at all is representable. Starting from
  11. 11 characters is easier to watch and makes no difference to the mechanism.
  12. 12 * Pre-tokenization is one simple rule (optional whitespace, then a run of
  13. 13 non-whitespace) rather than a long regex. That rule is why leading spaces
  14. 14 end up welded to words — which is the point worth seeing.
  15. 15 
  16. 16Tie-breaking: highest pair count wins; ties go to whichever pair was seen first
  17. 17in a left-to-right scan. First-seen order avoids comparing strings, which
  18. 18Python and JavaScript order differently for astral characters.
  19. 19"""
  20. 20 
  21. 21import re
  22. 22 
  23. 23PRETOKEN = re.compile(r"\s*\S+|\s+")
  24. 24 
  25. 25 
  26. 26def pretokenize(text):
  27. 27 """Split into chunks of optional leading whitespace plus a run of non-space."""
  28. 28 return PRETOKEN.findall(text)
  29. 29 
  30. 30 
  31. 31def word_counts(text):
  32. 32 """Unique chunks and how often each occurs, in first-seen order."""
  33. 33 counts = {}
  34. 34 for chunk in pretokenize(text):
  35. 35 counts[chunk] = counts.get(chunk, 0) + 1
  36. 36 return counts
  37. 37 
  38. 38 
  39. 39def _pair_stats(words):
  40. 40 """pair -> total count, in first-seen order (dicts preserve insertion order,
  41. 41 and so do JavaScript Maps — that is what makes tie-breaking identical)."""
  42. 42 counts = {}
  43. 43 for symbols, freq in words:
  44. 44 for i in range(len(symbols) - 1):
  45. 45 pair = (symbols[i], symbols[i + 1])
  46. 46 counts[pair] = counts.get(pair, 0) + freq
  47. 47 return counts
  48. 48 
  49. 49 
  50. 50def _merge_in(symbols, a, b):
  51. 51 """Replace every adjacent a,b with the joined symbol."""
  52. 52 out = []
  53. 53 i = 0
  54. 54 n = len(symbols)
  55. 55 while i < n:
  56. 56 if i < n - 1 and symbols[i] == a and symbols[i + 1] == b:
  57. 57 out.append(a + b)
  58. 58 i += 2
  59. 59 else:
  60. 60 out.append(symbols[i])
  61. 61 i += 1
  62. 62 return out
  63. 63 
  64. 64 
  65. 65def train(text, num_merges):
  66. 66 """Run BPE. Returns the merge list and a step-by-step record."""
  67. 67 words = [(list(chunk), freq) for chunk, freq in word_counts(text).items()]
  68. 68 
  69. 69 alphabet = {}
  70. 70 for chunk in word_counts(text):
  71. 71 for ch in chunk:
  72. 72 alphabet[ch] = True
  73. 73 
  74. 74 merges, steps = [], []
  75. 75 for _ in range(num_merges):
  76. 76 counts = _pair_stats(words)
  77. 77 if not counts:
  78. 78 break
  79. 79 # Strict >, scanning in first-seen order: the earliest of any tied pairs
  80. 80 # wins. No string comparison, so Python and JavaScript cannot disagree.
  81. 81 best = None
  82. 82 for pair, count in counts.items():
  83. 83 if best is None or count > counts[best]:
  84. 84 best = pair
  85. 85 if counts[best] < 2:
  86. 86 break # nothing repeats; merging is pointless
  87. 87 
  88. 88 a, b = best
  89. 89 words = [(_merge_in(sym, a, b), freq) for sym, freq in words]
  90. 90 merges.append((a, b))
  91. 91 steps.append({
  92. 92 "pair": [a, b],
  93. 93 "token": a + b,
  94. 94 "count": counts[best],
  95. 95 "vocab": len(alphabet) + len(merges),
  96. 96 })
  97. 97 
  98. 98 return {
  99. 99 "merges": merges,
  100. 100 "steps": steps,
  101. 101 "alphabet": sorted(alphabet),
  102. 102 "vocab_size": len(alphabet) + len(merges),
  103. 103 }
  104. 104 
  105. 105 
  106. 106def encode(word, merges):
  107. 107 """Apply a trained merge list to one chunk, in the order learned."""
  108. 108 symbols = list(word)
  109. 109 for a, b in merges:
  110. 110 symbols = _merge_in(symbols, a, b)
  111. 111 return symbols
  112. 112 
  113. 113 
  114. 114def encode_text(text, merges):
  115. 115 out = []
  116. 116 for chunk in pretokenize(text):
  117. 117 out.extend(encode(chunk, merges))
  118. 118 return out