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