spec100m / code /ngram.py
Akahsizrr's picture
Upload code/ngram.py with huggingface_hub
f015a7e verified
Raw History Blame Contribute Delete
5.46 kB
"""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)}'")