codekingpro/portable-devtools
114k
1import copy2from dataclasses import dataclass3 4import mmh35import numpy as np6from py_rust_stemmers import SnowballStemmer7 8from fastembed.common.utils import get_all_punctuation, remove_non_alphanumeric9from fastembed.sparse.sparse_embedding_base import SparseEmbedding10 11GAP = 3200012INT32_MAX = 2**31 - 113 14 15@dataclass16class WordEmbedding:17 word: str18 forms: list[str]19 count: int20 word_id: int21 embedding: list[float]22 23 24class SparseVectorConverter:25 def __init__(26 self,27 stopwords: set[str],28 stemmer: SnowballStemmer,29 k: float = 1.2,30 b: float = 0.75,31 avg_len: float = 150.0,32 ):33 punctuation = set(get_all_punctuation())34 special_tokens = {"[CLS]", "[SEP]", "[PAD]", "[UNK]", "[MASK]"}35 36 self.stemmer = stemmer37 self.unwanted_tokens = punctuation | special_tokens | stopwords38 39 self.k = k40 self.b = b41 self.avg_len = avg_len42 43 @classmethod44 def unkn_word_token_id(45 cls, word: str, shift: int46 ) -> int: # 2-3 words can collide in 1 index with this mapping, not considering mm3 collisions47 token_hash = abs(mmh3.hash(word))48 49 range_size = INT32_MAX - shift50 remapped_hash = shift + (token_hash % range_size)51 52 return remapped_hash53 54 def bm25_tf(self, num_occurrences: int, sentence_len: int) -> float:55 res = num_occurrences * (self.k + 1)56 res /= num_occurrences + self.k * (1 - self.b + self.b * sentence_len / self.avg_len)57 return res58 59 @classmethod60 def normalize_vector(cls, vector: list[float]) -> list[float]:61 norm = sum([x**2 for x in vector]) ** 0.562 if norm < 1e-8:63 return vector64 return [x / norm for x in vector]65 66 def clean_words(67 self, sentence_embedding: dict[str, WordEmbedding], token_max_length: int = 4068 ) -> dict[str, WordEmbedding]:69 """70 Clean miniCOIL-produced sentence_embedding, as unknown to the miniCOIL's stemmer tokens should fully resemble71 our BM25 token representation.72 73 sentence_embedding = {"9°": {"word": "9°", "word_id": -1, "count": 2, "embedding": [1], "forms": ["9°"]},74 "9": {"word": "9", "word_id": -1, "count": 2, "embedding": [1], "forms": ["9"]},75 "bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},76 "9°9": {"word": "9°9", "word_id": -1, "count": 1, "embedding": [1], "forms": ["9°9"]},77 "screech": {"word": "screech", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screech"]},78 "screeched": {"word": "screeched", "word_id": -1, "count": 1, "embedding": [1], "forms": ["screeched"]}79 }80 cleaned_embedding_ground_truth = {81 "9": {"word": "9", "word_id": -1, "count": 6, "embedding": [1], "forms": ["9°", "9", "9°9", "9°9"]},82 "bat": {"word": "bat", "word_id": 2, "count": 3, "embedding": [0.2, 0.1, -0.2, -0.2], "forms": ["bats", "bat"]},83 "screech": {"word": "screech", "word_id": -1, "count": 2, "embedding": [1], "forms": ["screech", "screeched"]}84 }85 """86 87 new_sentence_embedding: dict[str, WordEmbedding] = {}88 89 for word, embedding in sentence_embedding.items():90 # embedding = {91 # "word": "vector",92 # "forms": ["vector", "vectors"],93 # "count": 2,94 # "word_id": 1231,95 # "embedding": [0.1, 0.2, 0.3, 0.4]96 # }97 if embedding.word_id > 0:98 # Known word, no need to clean99 new_sentence_embedding[word] = embedding100 else:101 # Unknown word102 if word in self.unwanted_tokens:103 continue104 105 # Example complex word split:106 # word = `word^vec`107 word_cleaned = remove_non_alphanumeric(word).strip()108 # word_cleaned = `word vec`109 110 if len(word_cleaned) > 0:111 # Subwords: ['word', 'vec']112 for subword in word_cleaned.split():113 stemmed_subword: str = self.stemmer.stem_word(subword)114 if (115 len(stemmed_subword) <= token_max_length116 and stemmed_subword not in self.unwanted_tokens117 ):118 if stemmed_subword not in new_sentence_embedding:119 new_sentence_embedding[stemmed_subword] = copy.deepcopy(embedding)120 new_sentence_embedding[stemmed_subword].word = stemmed_subword121 else:122 new_sentence_embedding[stemmed_subword].count += embedding.count123 new_sentence_embedding[stemmed_subword].forms += embedding.forms124 125 return new_sentence_embedding126 127 def embedding_to_vector(128 self,129 sentence_embedding: dict[str, WordEmbedding],130 embedding_size: int,131 vocab_size: int,132 ) -> SparseEmbedding:133 """134 Convert miniCOIL sentence embedding to Qdrant sparse vector135 136 Example input:137 138 ```139 {140 "vector": WordEmbedding({ // Vocabulary word, encoded with miniCOIL normally141 "word": "vector",142 "forms": ["vector", "vectors"],143 "count": 2,144 "word_id": 1231,145 "embedding": [0.1, 0.2, 0.3, 0.4]146 }),147 "axiotic": WordEmbedding({ // Out-of-vocabulary word, fallback to BM25148 "word": "axiotic",149 "forms": ["axiotics"],150 "count": 1,151 "word_id": -1,152 })153 }154 ```155 156 """157 158 indices: list[int] = []159 values: list[float] = []160 161 # Example:162 # vocab_size = 10000163 # embedding_size = 4164 # GAP = 32000165 #166 # We want to start random words section from the bucket, that is guaranteed to not167 # include any vocab words.168 # We need (vocab_size * embedding_size) slots for vocab words.169 # Therefore we need (vocab_size * embedding_size) // GAP + 1 buckets for vocab words.170 # Therefore, we can start random words from bucket (vocab_size * embedding_size) // GAP + 1 + 1171 172 # ID at which the scope of OOV words starts173 unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP174 sentence_embedding_cleaned = self.clean_words(sentence_embedding)175 176 # Calculate sentence length after cleaning177 sentence_len = 0178 for embedding in sentence_embedding_cleaned.values():179 sentence_len += embedding.count180 181 for embedding in sentence_embedding_cleaned.values():182 word_id = embedding.word_id183 num_occurrences = embedding.count184 tf = self.bm25_tf(num_occurrences, sentence_len)185 if (186 word_id > 0187 ): # miniCOIL starts with ID 1, we generally won't have word_id == 0 (UNK), as we don't add188 # these words to sentence_embedding189 embedding_values = embedding.embedding190 normalized_embedding = self.normalize_vector(embedding_values)191 192 for val_id, value in enumerate(normalized_embedding):193 indices.append(194 word_id * embedding_size + val_id195 ) # since miniCOIL IDs start with 1196 values.append(value * tf)197 else:198 indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))199 values.append(tf)200 201 return SparseEmbedding(202 indices=np.array(indices, dtype=np.int32),203 values=np.array(values, dtype=np.float32),204 )205 206 def embedding_to_vector_query(207 self,208 sentence_embedding: dict[str, WordEmbedding],209 embedding_size: int,210 vocab_size: int,211 ) -> SparseEmbedding:212 """213 Same as `embedding_to_vector`, but no TF214 """215 216 indices: list[int] = []217 values: list[float] = []218 219 # ID at which the scope of OOV words starts220 unknown_words_shift = ((vocab_size * embedding_size) // GAP + 2) * GAP221 222 sentence_embedding_cleaned = self.clean_words(sentence_embedding)223 224 for embedding in sentence_embedding_cleaned.values():225 word_id = embedding.word_id226 tf = 1.0227 228 if word_id >= 0: # miniCOIL starts with ID 1229 embedding_values = embedding.embedding230 normalized_embedding = self.normalize_vector(embedding_values)231 232 for val_id, value in enumerate(normalized_embedding):233 indices.append(234 word_id * embedding_size + val_id235 ) # since miniCOIL IDs start with 1236 values.append(value * tf)237 else:238 indices.append(self.unkn_word_token_id(embedding.word, unknown_words_shift))239 values.append(tf)240 241 return SparseEmbedding(242 indices=np.array(indices, dtype=np.int32),243 values=np.array(values, dtype=np.float32),244 )245 