"""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