from pathlib import Path import json import heapq class SutraTokenizerFast: def __init__(self, directory): directory = Path(directory) with open(directory / "vocab.json", "r", encoding="utf-8") as f: self.vocab = json.load(f) with open(directory / "merges.json", "r", encoding="utf-8") as f: self.merges = [tuple(pair) for pair in json.load(f)] self.token_to_id = self.vocab # Invert vocab for fast decode lookup self.id_to_token = { int(idx): token for token, idx in self.vocab.items() } # Resolve special token IDs from vocabulary self.unk_id = self.vocab.get("<|unk|>") self.pad_token_id = self.vocab.get("<|pad|>") self.bos_token_id = self.vocab.get("<|bos|>") self.eos_token_id = self.vocab.get("<|eos|>") self.unk_token_id = self.unk_id # Set of special token IDs for filtering during decode self.special_token_ids = { self.pad_token_id, self.bos_token_id, self.eos_token_id, self.unk_token_id, } - {None} self.merge_rank = { pair: rank for rank, pair in enumerate(self.merges) } def encode(self, text): # ------------------------------------------------- # Normalize spaces # ------------------------------------------------- text = text.replace(" ", "▁") symbols = list(text) n = len(symbols) if n == 0: return [] if n == 1: return [ self.token_to_id.get( symbols[0], self.unk_id ) ] merge_rank = self.merge_rank token_to_id = self.token_to_id unk_id = self.unk_id # ------------------------------------------------- # Build initial heap # ------------------------------------------------- heap = [] for i in range(n - 1): pair = ( symbols[i], symbols[i + 1] ) rank = merge_rank.get(pair) if rank is not None: heapq.heappush( heap, (rank, i) ) # ------------------------------------------------- # Linked list # ------------------------------------------------- previous = [i - 1 for i in range(n)] following = [i + 1 for i in range(n)] following[-1] = -1 alive = bytearray(b"\x01") * n # ------------------------------------------------- # BPE merge loop # ------------------------------------------------- while heap: rank, left = heapq.heappop(heap) if not alive[left]: continue right = following[left] if right == -1 or not alive[right]: continue pair = ( symbols[left], symbols[right] ) if merge_rank.get(pair) != rank: continue # ------------------------------------------------- # Merge # ------------------------------------------------- symbols[left] += symbols[right] alive[right] = 0 next_index = following[right] following[left] = next_index if next_index != -1: previous[next_index] = left # ------------------------------------------------- # Previous pair # ------------------------------------------------- prev_index = previous[left] if prev_index != -1: new_pair = ( symbols[prev_index], symbols[left] ) new_rank = merge_rank.get(new_pair) if new_rank is not None: heapq.heappush( heap, ( new_rank, prev_index ) ) # ------------------------------------------------- # Next pair # ------------------------------------------------- if next_index != -1: new_pair = ( symbols[left], symbols[next_index] ) new_rank = merge_rank.get(new_pair) if new_rank is not None: heapq.heappush( heap, ( new_rank, left ) ) # ------------------------------------------------- # Convert surviving symbols to IDs # ------------------------------------------------- result = [] index = 0 while index != -1: if alive[index]: result.append( token_to_id.get( symbols[index], unk_id ) ) index = following[index] return result def decode(self, ids, skip_special_tokens=False, **kwargs): if skip_special_tokens: ids = [t for t in ids if t not in self.special_token_ids] text = "".join( self.id_to_token.get( int(idx), "<|unk|>" ) for idx in ids ) return text.replace("▁", " ")