File size: 4,569 Bytes
cbc7f31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Suffix automaton (SAM) drafting — SAM-Decoding style (arXiv:2411.10666).

Builds a suffix automaton over a static corpus + the generated text so far.
A query with the current context returns the longest suffix that appears in
the indexed text, and the tokens that followed it — i.e. a real, coherent
continuation drafted in O(match_len) with no model forward.

Why this beats fixed n-gram lookup:
  - Finds the LONGEST match (n-gram tables are capped at n).
  - Generalizes: any length from 2..len(context) is searchable.
  - The generated stream feeds back in, so self-consistent phrasing repeats
    cheaply (the model's own verified output becomes free tokens).

States:
  len[s]   : longest string length in state s
  link[s]  : suffix link
  next[s]  : dict token -> state (transitions)
  pos[s]   : position of first occurrence's END in the indexed text

Draft = text[pos[state] - match_len_in_state + 1 + k ...] i.e. whatever
followed the matched occurrence.
"""
from __future__ import annotations


class SuffixAutomaton:
    __slots__ = ("link", "length", "next", "pos", "last", "text")

    def __init__(self):
        self.link = [-1]
        self.length = [0]
        self.next = [{}]
        self.pos = [-1]
        self.last = 0
        self.text = []          # indexed token stream (corpus + generated)

    def extend(self, tok: int):
        """Append one token to the indexed text. O(1) amortized."""
        text = self.text
        text.append(tok)
        ppos = len(text) - 1

        cur = len(self.length)
        self.length.append(self.length[self.last] + 1)
        self.link.append(0)
        self.next.append({})
        self.pos.append(ppos)

        p = self.last
        while p != -1 and tok not in self.next[p]:
            self.next[p][tok] = cur
            p = self.link[p]
        if p == -1:
            self.link[cur] = 0
        else:
            q = self.next[p][tok]
            if self.length[p] + 1 == self.length[q]:
                self.link[cur] = q
            else:
                clone = len(self.length)
                self.length.append(self.length[p] + 1)
                self.link.append(self.link[q])
                self.next.append(dict(self.next[q]))
                self.pos.append(self.pos[q])
                while p != -1 and self.next[p].get(tok) == q:
                    self.next[p][tok] = clone
                    p = self.link[p]
                self.link[q] = self.link[cur] = clone
        self.last = cur

    def extend_many(self, toks):
        for t in toks:
            self.extend(t)

    def longest_match(self, context, min_len: int = 3):
        """Find the longest SUFFIX of `context` present in the indexed text.

        Walks the automaton over the whole context; at the end (v, l) is the
        longest suffix that occurs in the text. Returns (match_len, end_pos)
        where end_pos is the text index where the match ENDS, or (0, -1)."""
        v, l = 0, 0
        for tok in context:
            while v and tok not in self.next[v]:
                v = self.link[v]
                l = self.length[v]
            if tok in self.next[v]:
                v = self.next[v][tok]
                l += 1
            else:
                v, l = 0, 0
        if l < min_len:
            return 0, -1
        return l, self.pos[v]

    def draft(self, context, k: int, min_len: int = 4):
        """Return up to k continuation tokens after the longest context
        suffix match. If the best match's first occurrence sits at the end
        of the text (no continuation), fall back through suffix links to a
        shorter match with an available continuation."""
        v, l = 0, 0
        for tok in context:
            while v and tok not in self.next[v]:
                v = self.link[v]
                l = self.length[v]
            if tok in self.next[v]:
                v = self.next[v][tok]
                l += 1
            else:
                v, l = 0, 0
        # walk suffix links until a continuation exists
        n = len(self.text)
        while v and l >= min_len:
            end = self.pos[v]
            if 0 <= end < n - 1:
                out = self.text[end + 1: end + 1 + k]
                if out:
                    return out
            v = self.link[v]
            l = min(l, self.length[v])
        return None


def build_from_ids(ids, max_tokens: int | None = None):
    """Build a SAM over a token-id list (cap for memory)."""
    sam = SuffixAutomaton()
    sam.extend_many(ids[:max_tokens] if max_tokens else ids)
    return sam