Download tokenizer.py from KiranAN1988/NeemSutra-125M-Python-Instruct-Beta: direct link, hf CLI and curl.
- Browser
- Download file 5.96 kB
-
https://huggingface.co/KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/resolve/main/tokenizer.py
- Command line
-
hf download hf://KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/tokenizer.py
-
curl -L -o tokenizer.py https://huggingface.co/KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/resolve/main/tokenizer.py
5.96 kB
| 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("▁", " ") |