Download code/ngram.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 5.46 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/ngram.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/ngram.py
-
curl -L -o ngram.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/ngram.py
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)}'") | |