1010
1111logging .basicConfig (level = logging .INFO , format = "%(asctime)s - %(levelname)s - %(message)s" )
1212
13+ _GLOBAL_CACHES = {}
14+ _GLOBAL_LOCKS = {}
15+
1316
1417class EmbeddingModel (ABC ):
1518 def __init__ (self , model_name : str , cache : bool = False , cache_size : int = 10000 ):
@@ -20,8 +23,19 @@ def __init__(self, model_name: str, cache: bool = False, cache_size: int = 10000
2023 :param cache_size: Maximum cache size when caching is enabled, defaults to 10,000.
2124 """
2225 self .model_name = model_name
23- self ._cache = LRUCache (maxsize = cache_size ) if cache else None
24- self ._cache_lock = Lock () if cache else None
26+ self .use_cache = cache
27+
28+ if self .use_cache :
29+ if model_name not in _GLOBAL_CACHES :
30+ _GLOBAL_CACHES [model_name ] = LRUCache (maxsize = cache_size )
31+ _GLOBAL_LOCKS [model_name ] = Lock ()
32+
33+ self ._cache = _GLOBAL_CACHES [model_name ]
34+ self ._cache_lock = _GLOBAL_LOCKS [model_name ]
35+
36+ else :
37+ self ._cache = None
38+ self ._cache_lock = None
2539
2640 @abstractmethod
2741 def get_embedding (self , text : str ) -> Sequence [float ]:
@@ -47,22 +61,31 @@ def add_to_cache(self, text: str, embedding: Sequence[float]):
4761 :param text: The input text to be cached.
4862 :param embedding: The embedding of the input text.
4963 """
50- if self ._cache_lock and self ._cache is not None :
64+ if self ._cache_lock is not None and self ._cache is not None :
5165 with self ._cache_lock :
5266 self ._cache [text ] = embedding
5367
68+ def add_batch_to_cache (self , texts : Sequence [str ], embeddings : Sequence [Sequence [float ]]):
69+ """Acquire lock once for the entire batch update."""
70+ if self ._cache_lock is not None and self ._cache is not None :
71+ with self ._cache_lock :
72+ for text , emb in zip (texts , embeddings ):
73+ self ._cache [text ] = emb
74+
5475 def get_from_cache (self , text : str ) -> Optional [Sequence [float ]]:
5576 """Retrieve an embedding from the cache.
5677
5778 :param text: Cached input text.
5879 :return: Embedding of the cached input text or `None` if not present.
5980 """
60- if self ._cache_lock and self ._cache is not None :
81+ if self ._cache_lock is not None and self ._cache is not None :
6182 with self ._cache_lock :
6283 return self ._cache .get (text , None )
6384 return None
6485
65- def get_cached_embeddings (self , messages : List [str ]) -> Tuple [List [Sequence [float ]], List [int ], List [str ]]:
86+ def get_cached_embeddings (
87+ self , messages : List [str ]
88+ ) -> Tuple [List [Optional [Sequence [float ]]], List [int ], List [str ]]:
6689 """Retrieve cached embeddings and identify uncached messages.
6790
6891 :param messages: A list of input text messages.
@@ -71,16 +94,22 @@ def get_cached_embeddings(self, messages: List[str]) -> Tuple[List[Sequence[floa
7194 - A list of indices for uncached messages.
7295 - A list of uncached messages.
7396 """
74- embeddings , uncached_indices , uncached_messages = [], [], []
97+ if self ._cache_lock is None or self ._cache is None :
98+ empty_embeddings : List [Optional [Sequence [float ]]] = [None ] * len (messages )
99+ return empty_embeddings , list (range (len (messages ))), messages
75100
76- for i , msg in enumerate (messages ):
77- cached = self .get_from_cache (msg )
78- if cached :
101+ embeddings : List [Optional [Sequence [float ]]] = []
102+ uncached_indices : List [int ] = []
103+ uncached_messages : List [str ] = []
104+
105+ with self ._cache_lock :
106+ for i , msg in enumerate (messages ):
107+ cached = self ._cache .get (msg , None )
79108 embeddings .append (cached )
80- else :
81- embeddings .append (None )
82- uncached_indices .append (i )
83- uncached_messages . append ( msg )
109+ if cached is None :
110+ uncached_indices .append (i )
111+ uncached_messages .append (msg )
112+
84113 return embeddings , uncached_indices , uncached_messages
85114
86115 def sanitize (self , message : str ) -> str :
@@ -120,8 +149,7 @@ def get_embedding(self, text: str) -> Sequence[float]:
120149 return []
121150 text = self .sanitize (text )
122151
123- if self ._cache :
124- # Check cache
152+ if self ._cache is not None :
125153 cached = self .get_from_cache (text )
126154 if cached :
127155 return cached
@@ -134,7 +162,7 @@ def get_embedding(self, text: str) -> Sequence[float]:
134162 return embedding
135163 except Exception as e :
136164 logging .error (f"Error getting embedding for { text } : { e } " )
137- return []
165+ raise
138166
139167 def get_embeddings (self , messages : List [str ]) -> Sequence [Sequence [float ]]:
140168 """Retrieve embeddings for a list of text messages.
@@ -145,19 +173,19 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
145173
146174 sanitized_messages = [self .sanitize (msg ) for msg in messages ]
147175
148- if self ._cache :
176+ if self ._cache is not None :
149177 embeddings , uncached_indices , uncached_messages = self .get_cached_embeddings (sanitized_messages )
150178
151179 if uncached_messages :
152180 try :
153181 response = openai .embeddings .create (model = self .model_name , input = uncached_messages )
154182 new_embeddings = [item .embedding for item in response .data ]
183+ self .add_batch_to_cache (uncached_messages , new_embeddings )
155184 for idx , embedding in zip (uncached_indices , new_embeddings ):
156- self .add_to_cache (sanitized_messages [idx ], embedding )
157185 embeddings [idx ] = embedding
158186 except Exception as e :
159187 logging .error (f"Error in processing chunk: { e } " )
160- return []
188+ raise
161189
162190 return [emb for emb in embeddings if emb is not None ]
163191
@@ -167,12 +195,12 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
167195 return embeddings
168196 except Exception as e :
169197 logging .error (f"Failed processing messages: { e } " )
170- return []
198+ raise
171199
172200
173201class HuggingFaceAdapter (EmbeddingModel ):
174202 _model_cache = {}
175- _load_count = 0 # For testing
203+ _load_count = 0 # For testing
176204
177205 def __init__ (self , model_name : str = "sentence-transformers/all-MiniLM-L6-v2" , cache : bool = False ):
178206 """Initialize the Hugging Face adapter with a specified model name.
@@ -200,7 +228,7 @@ def get_embedding(self, text: str) -> Sequence[float]:
200228 text = self .sanitize (text )
201229
202230 # Check cache
203- if self ._cache :
231+ if self ._cache is not None :
204232 cached = self .get_from_cache (text )
205233 if cached :
206234 return cached
@@ -212,7 +240,7 @@ def get_embedding(self, text: str) -> Sequence[float]:
212240 return embedding
213241 except Exception as e :
214242 logging .error (f"Error getting embedding for { text } : { e } " )
215- return []
243+ raise
216244
217245 def get_embeddings (self , messages : List [str ]) -> Sequence [Sequence [float ]]:
218246 """Retrieve embeddings for a list of text messages using MPNet.
@@ -221,7 +249,7 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
221249 :return: A sequence of embedding vectors.
222250 """
223251 sanitized_messages = [self .sanitize (msg ) for msg in messages ]
224- if self ._cache :
252+ if self ._cache is not None :
225253 embeddings , uncached_indices , uncached_messages = self .get_cached_embeddings (sanitized_messages )
226254
227255 if uncached_messages :
@@ -230,12 +258,12 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
230258 flattened_embeddings = [
231259 [float (element ) for element in row ] for row in new_embeddings if row is not None
232260 ]
261+ self .add_batch_to_cache (uncached_messages , flattened_embeddings )
233262 for idx , embedding in zip (uncached_indices , flattened_embeddings ):
234- self .add_to_cache (sanitized_messages [idx ], embedding )
235263 embeddings [idx ] = embedding
236264 except Exception as e :
237265 logging .error (f"Failed processing messages: { e } " )
238- return []
266+ raise
239267
240268 return [emb for emb in embeddings if emb is not None ]
241269
@@ -245,7 +273,7 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
245273 return flattened_embeddings
246274 except Exception as e :
247275 logging .error (f"Failed processing messages: { e } " )
248- return []
276+ raise
249277
250278
251279class OllamaAdapter (EmbeddingModel ):
@@ -260,7 +288,7 @@ def get_embedding(self, text: str) -> Sequence[float]:
260288 logging .warning ("Empty text passed to get_embedding" )
261289 return []
262290 text = self .sanitize (text )
263- if self ._cache :
291+ if self ._cache is not None :
264292 cached = self .get_from_cache (text )
265293 if cached :
266294 return cached
@@ -270,31 +298,31 @@ def get_embedding(self, text: str) -> Sequence[float]:
270298 return embedding
271299 except Exception as e :
272300 logging .error (f"Error getting embedding for { text } : { e } " )
273- return []
301+ raise
274302
275303 def get_embeddings (self , messages : List [str ]) -> Sequence [Sequence [float ]]:
276304 sanitized_messages = [self .sanitize (msg ) for msg in messages ]
277305
278- if self ._cache :
306+ if self ._cache is not None :
279307 embeddings , uncached_indices , uncached_messages = self .get_cached_embeddings (sanitized_messages )
280308
281309 if uncached_messages :
282310 try :
283311 new_embeddings = self .client .embed (self .model_name , uncached_messages ).get ("embeddings" )
312+ self .add_batch_to_cache (uncached_messages , new_embeddings )
284313 for idx , embedding in zip (uncached_indices , new_embeddings ):
285- self .add_to_cache (sanitized_messages [idx ], embedding )
286314 embeddings [idx ] = embedding
287315 except Exception as e :
288316 logging .error (f"Failed processing messages: { e } " )
289- return []
317+ raise
290318
291319 return [emb for emb in embeddings if emb is not None ]
292320
293321 try :
294322 return self .client .embed (self .model_name , sanitized_messages ).get ("embeddings" )
295323 except Exception as e :
296324 logging .error (f"Failed processing messages: { e } " )
297- return []
325+ raise
298326
299327
300328class Vectorizer :
0 commit comments