sweedworks

← all sources

ngram.py

The n-gram model and sampling knobs behind /predict/

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

  1. 1"""A small n-gram language model, and the sampling knobs that sit on top of it.
  2. 2 
  3. 3The point is not that this is a good model — it is a lookup table of what
  4. 4followed what. The point is that every step between "a distribution over next
  5. 5tokens" and "a word on your screen" is arithmetic you can watch, and those steps
  6. 6are identical in a real model. Temperature, top-k and top-p do exactly this to
  7. 7a 200,000-way distribution instead of a twenty-way one.
  8. 8 
  9. 9Mirrored by predict/ngram.js; checkngram.py fails the build if they diverge.
  10. 10Sampling uses a seeded PRNG (mulberry32) so that a given seed produces the same
  11. 11text in both implementations, and for the reader, every time.
  12. 12"""
  13. 13 
  14. 14from bpe import pretokenize
  15. 15 
  16. 16 
  17. 17def tokenize(text):
  18. 18 """Same rule as the BPE page: optional leading space, then non-space."""
  19. 19 return pretokenize(text)
  20. 20 
  21. 21 
  22. 22def train(text, order):
  23. 23 """context tuple -> {next token: count}, for every context length <= order."""
  24. 24 tokens = tokenize(text)
  25. 25 model = {}
  26. 26 for k in range(order + 1):
  27. 27 table = {}
  28. 28 for i in range(len(tokens) - k):
  29. 29 ctx = tuple(tokens[i:i + k])
  30. 30 nxt = tokens[i + k]
  31. 31 table.setdefault(ctx, {})
  32. 32 table[ctx][nxt] = table[ctx].get(nxt, 0) + 1
  33. 33 model[k] = table
  34. 34 # First-appearance order, not sorted: JavaScript orders strings by UTF-16
  35. 35 # code unit and Python by code point, and they disagree on astral characters.
  36. 36 vocab, seen = [], set()
  37. 37 for t in tokens:
  38. 38 if t not in seen:
  39. 39 seen.add(t)
  40. 40 vocab.append(t)
  41. 41 return {"model": model, "order": order, "tokens": tokens, "vocab": vocab}
  42. 42 
  43. 43 
  44. 44def distribution(trained, context):
  45. 45 """Probabilities for the next token, backing off to shorter contexts.
  46. 46 
  47. 47 Returns (list of (token, probability) sorted by probability desc, order used).
  48. 48 """
  49. 49 order = trained["order"]
  50. 50 ctx = tuple(context)
  51. 51 for k in range(min(order, len(ctx)), -1, -1):
  52. 52 table = trained["model"][k]
  53. 53 key = ctx[len(ctx) - k:] if k else ()
  54. 54 if key in table:
  55. 55 counts = table[key]
  56. 56 total = sum(counts.values())
  57. 57 items = [(t, c / total) for t, c in counts.items()]
  58. 58 # Sort by probability, then by first appearance in the vocabulary,
  59. 59 # so ties resolve identically in both implementations.
  60. 60 order_index = {t: i for i, t in enumerate(trained["vocab"])}
  61. 61 items.sort(key=lambda p: (-p[1], order_index[p[0]]))
  62. 62 return items, k
  63. 63 return [], 0
  64. 64 
  65. 65 
  66. 66def _total(values):
  67. 67 """Naive left-to-right float accumulation.
  68. 68 
  69. 69 Deliberately not sum(): since 3.12 CPython uses Neumaier compensated
  70. 70 summation for floats, so sum() returns 1.0 where JavaScript's += loop
  71. 71 returns 0.9999999999999999. Being more accurate than the browser is still
  72. 72 being different from it, and that difference moved where top-p cut the
  73. 73 distribution. Matching the browser is the requirement here.
  74. 74 """
  75. 75 total = 0.0
  76. 76 for v in values:
  77. 77 total += v
  78. 78 return total
  79. 79 
  80. 80 
  81. 81def apply_temperature(items, temp):
  82. 82 """p -> p^(1/T), renormalised. T<1 sharpens, T>1 flattens, T→0 is greedy."""
  83. 83 if temp <= 0:
  84. 84 # The limit: all mass on the most likely token.
  85. 85 return [(t, 1.0 if i == 0 else 0.0) for i, (t, _) in enumerate(items)]
  86. 86 scaled = [(t, p ** (1.0 / temp)) for t, p in items]
  87. 87 total = _total(p for _, p in scaled)
  88. 88 if total == 0:
  89. 89 return items
  90. 90 return [(t, p / total) for t, p in scaled]
  91. 91 
  92. 92 
  93. 93def apply_top_k(items, k):
  94. 94 """Keep the k most likely tokens, renormalise. 0 disables."""
  95. 95 if k <= 0 or k >= len(items):
  96. 96 return items
  97. 97 kept = items[:k]
  98. 98 total = _total(p for _, p in kept)
  99. 99 return [(t, p / total) for t, p in kept]
  100. 100 
  101. 101 
  102. 102def apply_top_p(items, p_threshold):
  103. 103 """Nucleus: keep the smallest set whose probability sums past the threshold."""
  104. 104 if p_threshold <= 0 or p_threshold >= 1:
  105. 105 return items
  106. 106 kept, running = [], 0.0
  107. 107 for token, p in items:
  108. 108 kept.append((token, p))
  109. 109 running += p
  110. 110 if running >= p_threshold:
  111. 111 break
  112. 112 total = _total(p for _, p in kept)
  113. 113 return [(t, p / total) for t, p in kept]
  114. 114 
  115. 115 
  116. 116class Rng:
  117. 117 """mulberry32 — small, seedable, and identical in JavaScript."""
  118. 118 
  119. 119 def __init__(self, seed):
  120. 120 self.state = seed & 0xFFFFFFFF
  121. 121 
  122. 122 def next(self):
  123. 123 self.state = (self.state + 0x6D2B79F5) & 0xFFFFFFFF
  124. 124 t = self.state
  125. 125 t = ((t ^ (t >> 15)) * (t | 1)) & 0xFFFFFFFF
  126. 126 t = (t ^ (t + ((t ^ (t >> 7)) * (t | 61) & 0xFFFFFFFF))) & 0xFFFFFFFF
  127. 127 return ((t ^ (t >> 14)) & 0xFFFFFFFF) / 4294967296.0
  128. 128 
  129. 129 
  130. 130def pick(items, rng):
  131. 131 """Sample one token from a normalised distribution."""
  132. 132 if not items:
  133. 133 return None
  134. 134 r = rng.next()
  135. 135 acc = 0.0
  136. 136 for token, p in items:
  137. 137 acc += p
  138. 138 if r < acc:
  139. 139 return token
  140. 140 return items[-1][0]
  141. 141 
  142. 142 
  143. 143def shape(items, temp=1.0, top_k=0, top_p=0.0):
  144. 144 """The full pipeline, in the order a real sampler applies it."""
  145. 145 out = apply_temperature(items, temp)
  146. 146 out = apply_top_k(out, top_k)
  147. 147 out = apply_top_p(out, top_p)
  148. 148 return out
  149. 149 
  150. 150 
  151. 151def generate(trained, prompt, n, temp=1.0, top_k=0, top_p=0.0, seed=7):
  152. 152 """Produce n tokens, returning both the text and a record of each choice."""
  153. 153 rng = Rng(seed)
  154. 154 context = list(tokenize(prompt))
  155. 155 produced, trace = [], []
  156. 156 for _ in range(n):
  157. 157 items, used = distribution(trained, context)
  158. 158 if not items:
  159. 159 break
  160. 160 shaped = shape(items, temp, top_k, top_p)
  161. 161 token = pick(shaped, rng)
  162. 162 if token is None:
  163. 163 break
  164. 164 chosen_p = dict(shaped).get(token, 0.0)
  165. 165 trace.append({
  166. 166 "token": token,
  167. 167 "p": chosen_p,
  168. 168 "considered": len(shaped),
  169. 169 "of": len(items),
  170. 170 "context_order": used,
  171. 171 })
  172. 172 produced.append(token)
  173. 173 context.append(token)
  174. 174 return {"text": "".join(produced), "tokens": produced, "trace": trace}