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