"""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)}'")