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