Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
vocab_resolver.py203 linesDownload Raw Back to utils
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 
codekingpro/portable-devtools · Team Ai