KiranAN1988's picture
Initial beta release
2e30a77 verified
Raw History Blame Contribute Delete
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("▁", " ")