codekingpro/portable-devtools
114k
1from collections import defaultdict2from typing import Iterable3 4from py_rust_stemmers import SnowballStemmer5import numpy as np6from tokenizers import Tokenizer7from numpy.typing import NDArray8 9from fastembed.common.types import NumpyArray10 11 12class VocabTokenizerBase:13 def tokenize(self, sentence: str) -> NumpyArray:14 raise NotImplementedError()15 16 def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:17 raise NotImplementedError()18 19 20class VocabTokenizer(VocabTokenizerBase):21 def __init__(self, tokenizer: Tokenizer):22 self.tokenizer = tokenizer23 24 def tokenize(self, sentence: str) -> NumpyArray:25 return np.array(self.tokenizer.encode(sentence).ids)26 27 def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:28 return [self.tokenizer.id_to_token(token_id) for token_id in token_ids]29 30 31class VocabResolver:32 def __init__(self, tokenizer: VocabTokenizerBase, stopwords: set[str], stemmer: SnowballStemmer):33 # Word to id mapping34 self.vocab: dict[str, int] = {}35 # Id to word mapping36 self.words: list[str] = []37 # Lemma to word mapping38 self.stem_mapping: dict[str, str] = {}39 self.tokenizer: VocabTokenizerBase = tokenizer40 self.stemmer = stemmer41 self.stopwords: set[str] = stopwords42 43 def tokenize(self, sentence: str) -> NumpyArray:44 return self.tokenizer.tokenize(sentence)45 46 def lookup_word(self, word_id: int) -> str:47 if word_id == 0:48 return "UNK"49 return self.words[word_id - 1]50 51 def convert_ids_to_tokens(self, token_ids: NumpyArray) -> list[str]:52 return self.tokenizer.convert_ids_to_tokens(token_ids)53 54 def vocab_size(self) -> int:55 # We need +1 for UNK token56 return len(self.vocab) + 157 58 def save_vocab(self, path: str) -> None:59 with open(path, "w") as f:60 for word in self.words:61 f.write(word + "\n")62 63 def save_json_vocab(self, path: str) -> None:64 import json65 66 with open(path, "w") as f:67 json.dump({"vocab": self.words, "stem_mapping": self.stem_mapping}, f, indent=2)68 69 def load_json_vocab(self, path: str) -> None:70 import json71 72 with open(path, "r") as f:73 data = json.load(f)74 self.words = data["vocab"]75 self.vocab = {word: idx + 1 for idx, word in enumerate(self.words)}76 self.stem_mapping = data["stem_mapping"]77 78 def add_word(self, word: str) -> None:79 if word not in self.vocab:80 self.vocab[word] = len(self.vocab) + 181 self.words.append(word)82 stem = self.stemmer.stem_word(word)83 if stem not in self.stem_mapping:84 self.stem_mapping[stem] = word85 else:86 existing_word = self.stem_mapping[stem]87 if len(existing_word) > len(word):88 # Prefer shorter words for the same stem89 # Example: "swim" is preferred over "swimming"90 self.stem_mapping[stem] = word91 92 def load_vocab(self, path: str) -> None:93 with open(path, "r") as f:94 for line in f:95 self.add_word(line.strip())96 97 @classmethod98 def _reconstruct_bpe(99 cls, bpe_tokens: Iterable[tuple[int, str]]100 ) -> list[tuple[str, list[int]]]:101 result: list[tuple[str, list[int]]] = []102 acc: str = ""103 acc_idx: list[int] = []104 105 continuing_subword_prefix = "##"106 continuing_subword_prefix_len = len(continuing_subword_prefix)107 108 for idx, token in bpe_tokens:109 if token.startswith(continuing_subword_prefix):110 acc += token[continuing_subword_prefix_len:]111 acc_idx.append(idx)112 else:113 if acc:114 result.append((acc, acc_idx))115 acc_idx = []116 acc = token117 acc_idx.append(idx)118 119 if acc:120 result.append((acc, acc_idx))121 return result122 123 def resolve_tokens(124 self, token_ids: NDArray[np.int64]125 ) -> tuple[NDArray[np.int64], dict[int, int], dict[str, int], dict[str, list[str]]]:126 """127 Mark known tokens (including composed tokens) with vocab ids.128 129 Args:130 token_ids: (seq_len) - list of ids of tokens131 Example:132 [133 101, 3897, 19332, 12718, 23348,134 1010, 1996, 7151, 2296, 4845,135 2359, 2005, 4234, 1010, 4332,136 2871, 3191, 2062, 102137 ]138 139 returns:140 - token_ids with vocab ids141 [142 0, 151, 151, 0, 0,143 912, 0, 0, 0, 332,144 332, 332, 0, 7121, 191,145 0, 0, 332, 0146 ]147 - counts of each token148 {149 151: 1,150 332: 3,151 7121: 1,152 191: 1,153 912: 1154 }155 - oov counts of each token156 {157 "the": 1,158 "a": 1,159 "[CLS]": 1,160 "[SEP]": 1,161 ...162 }163 - forms of each token164 {165 "hello": ["hello"],166 "world": ["worlds", "world", "worlding"],167 }168 169 """170 tokens = self.convert_ids_to_tokens(token_ids)171 tokens_mapping = self._reconstruct_bpe(enumerate(tokens))172 173 counts: dict[int, int] = defaultdict(int)174 oov_count: dict[str, int] = defaultdict(int)175 176 forms: dict[str, list[str]] = defaultdict(list)177 178 for token, mapped_token_ids in tokens_mapping:179 vocab_id = 0180 if token in self.stopwords:181 vocab_id = 0182 elif token in self.vocab:183 vocab_id = self.vocab[token]184 forms[token].append(token)185 elif token in self.stem_mapping:186 vocab_id = self.vocab[self.stem_mapping[token]]187 forms[self.stem_mapping[token]].append(token)188 else:189 stem = self.stemmer.stem_word(token)190 if stem in self.stem_mapping:191 vocab_id = self.vocab[self.stem_mapping[stem]]192 forms[self.stem_mapping[stem]].append(token)193 194 for token_id in mapped_token_ids:195 token_ids[token_id] = vocab_id196 197 if vocab_id == 0:198 oov_count[token] += 1199 else:200 counts[vocab_id] += 1201 return token_ids, counts, oov_count, forms202 203 