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
|