File size: 5,462 Bytes
f015a7e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | """N-gram model for free draft token extension.
The n-gram is a statistical lookup table: given the last n tokens, predict
the most likely next token from training data. It costs ZERO neural network
params and ~1μs per lookup (Python dict).
Used in V5Engine to extend each speculative step:
- Transformer + Medusa: K+1 tokens per forward pass (costs ~20ms)
- N-gram extension: +M tokens per step (costs ~0.5ms, just dict lookups)
- Effective throughput: (K+1+M) / step_time
The n-gram also improves quality: it predicts tokens based on real observed
patterns instead of random untrained Medusa guesses. On real text, a 5-gram
model captures common phrases, punctuation patterns, and local syntax.
Storage:
- n=5, vocab=50257: key = tuple of 5 token ids, value = most likely next token
- On 100M tokens of training data: ~5M unique 5-grams, ~200 MB dict
- Compact mode: store only the argmax next token (not full distribution)
"""
from __future__ import annotations
import os
import pickle
from collections import defaultdict, Counter
from typing import Optional
class NGramModel:
"""N-gram model with fallback to (n-1)-gram, ..., down to unigram.
Stores the most likely next token for each observed n-gram.
Falls back to shorter context when the full n-gram is unseen.
"""
def __init__(self, n: int = 5, vocab_size: int = 50257):
self.n = n
self.vocab_size = vocab_size
# tables[k] = dict: tuple of k tokens -> most likely next token
# k ranges from 1 (unigram) to n (full n-gram)
self.tables: list[dict[tuple, int]] = [dict() for _ in range(n + 1)]
self._total_tokens = 0
def train(self, token_ids: list[int], verbose: bool = True):
"""Build n-gram tables from a list of token ids.
For each context length k (1..n), count (k-gram, next_token) pairs
and store the most frequent next token for each k-gram.
"""
self._total_tokens = len(token_ids)
if verbose:
print(f" ngram: training on {self._total_tokens:,} tokens, n={self.n}")
for k in range(1, self.n + 1):
counts: dict[tuple, Counter] = defaultdict(Counter)
for i in range(len(token_ids) - k):
context = tuple(token_ids[i:i + k])
next_tok = token_ids[i + k]
counts[context][next_tok] += 1
# Store only the argmax next token for each context
self.tables[k] = {ctx: c.most_common(1)[0][0]
for ctx, c in counts.items()}
if verbose:
print(f" {k}-gram: {len(self.tables[k]):,} entries")
def predict(self, context: list[int]) -> Optional[int]:
"""Predict the next token given a context of recent tokens.
Falls back from n-gram to (n-1)-gram ... down to unigram.
Returns None if no prediction is available.
"""
for k in range(min(self.n, len(context)), 0, -1):
key = tuple(context[-k:])
if key in self.tables[k]:
return self.tables[k][key]
return None
def predict_batch(self, context: list[int], n_tokens: int) -> list[int]:
"""Chain-predict n_tokens, feeding each prediction back as context.
Stops early if a prediction fails (returns None).
Returns the predicted tokens (may be shorter than n_tokens).
"""
out = []
ctx = list(context)
for _ in range(n_tokens):
pred = self.predict(ctx)
if pred is None:
break
out.append(pred)
ctx.append(pred)
return out
def save(self, path: str):
"""Save n-gram tables to disk."""
with open(path, "wb") as f:
pickle.dump({
"n": self.n,
"vocab_size": self.vocab_size,
"tables": self.tables,
"total_tokens": self._total_tokens,
}, f)
size_mb = os.path.getsize(path) / 1e6
print(f" ngram: saved to {path} ({size_mb:.1f} MB)")
def load(self, path: str):
"""Load n-gram tables from disk."""
with open(path, "rb") as f:
data = pickle.load(f)
self.n = data["n"]
self.vocab_size = data["vocab_size"]
self.tables = data["tables"]
self._total_tokens = data["total_tokens"]
print(f" ngram: loaded from {path} (n={self.n}, "
f"{sum(len(t) for t in self.tables):,} entries)")
def stats(self) -> dict:
return {
"n": self.n,
"entries": sum(len(t) for t in self.tables),
"total_tokens": self._total_tokens,
"size_estimate_mb": sum(len(t) for t in self.tables) * 0.05,
}
if __name__ == "__main__":
# Quick test
import tiktoken
enc = tiktoken.get_encoding("gpt2")
text = ("The quick brown fox jumps over the lazy dog. " * 1000 +
"To be or not to be, that is the question. " * 1000)
tokens = enc.encode_ordinary(text)
print(f"Training text: {len(tokens)} tokens")
ng = NGramModel(n=5)
ng.train(tokens)
# Test prediction
test = "The quick brown"
test_ids = enc.encode_ordinary(test)
pred = ng.predict(test_ids)
print(f"Context: '{test}' -> predicted: '{enc.decode([pred])}'")
# Test batch prediction
batch = ng.predict_batch(test_ids, 10)
print(f"Batch (10): '{enc.decode(batch)}'")
|