Skip to content

Commit f7d9465

Browse files
authored
Merge pull request #182 from SCAI-BIO/refactor/embedding-cache-optimization
Refactor/embedding cache optimization
2 parents ffcad8b + af6ede9 commit f7d9465

2 files changed

Lines changed: 104 additions & 74 deletions

File tree

datastew/embedding.py

Lines changed: 61 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,9 @@
1010

1111
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
1212

13+
_GLOBAL_CACHES = {}
14+
_GLOBAL_LOCKS = {}
15+
1316

1417
class 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

173201
class 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

251279
class 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

300328
class Vectorizer:

tests/test_embedding.py

Lines changed: 43 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -1,82 +1,84 @@
11
import unittest
2-
from time import time
32
from typing import Sequence
3+
from unittest.mock import patch
44

5-
from datastew.embedding import HuggingFaceAdapter
5+
from datastew.embedding import _GLOBAL_CACHES, HuggingFaceAdapter
66

77

88
class TestEmbedding(unittest.TestCase):
99

1010
def setUp(self):
11-
self.hugging_face_adapter = HuggingFaceAdapter(cache=True)
11+
_GLOBAL_CACHES.clear()
12+
self.adapter = HuggingFaceAdapter(cache=True)
1213

1314
def test_hugging_face_adapter_get_embedding(self):
1415
text = "This is a test sentence."
15-
embedding = self.hugging_face_adapter.get_embedding(text)
16+
embedding = self.adapter.get_embedding(text)
1617
self.assertIsInstance(embedding, Sequence)
1718

1819
def test_hugging_face_adapter_get_embeddings(self):
1920
messages = [f"This is message {i}." for i in range(20)]
20-
embeddings = self.hugging_face_adapter.get_embeddings(messages)
21+
embeddings = self.adapter.get_embeddings(messages)
2122
self.assertIsInstance(embeddings, Sequence)
2223
self.assertEqual(len(embeddings), len(messages))
2324

2425
def test_sanitization(self):
2526
text1 = " Test"
2627
text2 = "test "
27-
embedding1 = self.hugging_face_adapter.get_embedding(text1)
28-
embedding2 = self.hugging_face_adapter.get_embedding(text2)
28+
embedding1 = self.adapter.get_embedding(text1)
29+
embedding2 = self.adapter.get_embedding(text2)
2930
self.assertSequenceEqual(embedding1, embedding2)
3031

31-
def test_caching_get_embedding(self):
32-
text = "This is a test sentence."
33-
if self.hugging_face_adapter._cache:
34-
self.hugging_face_adapter._cache.clear()
32+
def test_caching_get_embedding_deterministic(self):
33+
text = "deterministic test sentence"
3534

36-
# Measure time for the first call
37-
start_time = time()
38-
embedding1 = self.hugging_face_adapter.get_embedding(text)
39-
first_call_time = time() - start_time
35+
with patch.object(self.adapter.model, "encode", wraps=self.adapter.model.encode) as spy_encode:
36+
emb1 = self.adapter.get_embedding(text)
37+
spy_encode.assert_called_once()
4038

41-
# Measure time for the second call
42-
start_time = time()
43-
embedding2 = self.hugging_face_adapter.get_embedding(text)
44-
second_call_time = time() - start_time
39+
emb2 = self.adapter.get_embedding(text)
40+
self.assertEqual(spy_encode.call_count, 1)
41+
self.assertSequenceEqual(emb1, emb2)
4542

46-
self.assertLess(second_call_time, first_call_time)
47-
self.assertSequenceEqual(embedding1, embedding2)
43+
def test_caching_get_embeddings_batch(self):
44+
messages = [f"Batch message {i}." for i in range(5)]
4845

49-
def test_caching_get_embeddings(self):
50-
messages = [f"This is message {i}." for i in range(20)]
51-
if self.hugging_face_adapter._cache:
52-
self.hugging_face_adapter._cache.clear()
46+
with patch.object(self.adapter.model, "encode", wraps=self.adapter.model.encode) as spy_encode:
47+
embeddings1 = self.adapter.get_embeddings(messages)
48+
self.assertEqual(spy_encode.call_count, 1)
49+
embeddings2 = self.adapter.get_embeddings(messages)
50+
self.assertEqual(spy_encode.call_count, 1)
5351

54-
start_time = time()
55-
embeddings1 = self.hugging_face_adapter.get_embeddings(messages)
56-
first_call_time = time() - start_time
52+
for emb1, emb2 in zip(embeddings1, embeddings2):
53+
self.assertSequenceEqual(emb1, emb2)
5754

58-
start_time = time()
59-
embeddings2 = self.hugging_face_adapter.get_embeddings(messages)
60-
second_call_time = time() - start_time
55+
def test_partial_cache_hit_batch(self):
56+
messages = ["A", "B", "C"]
57+
# Cache "A" manually
58+
self.adapter.add_batch_to_cache([self.adapter.sanitize("A")], [[0.1, 0.2]])
6159

62-
self.assertLess(second_call_time, first_call_time)
60+
with patch.object(self.adapter.model, "encode", wraps=self.adapter.model.encode) as spy_encode:
61+
embeddings = self.adapter.get_embeddings(messages)
62+
# encode should only be called once, and only for ["b", "c"]
63+
spy_encode.assert_called_once_with(["b", "c"], show_progress_bar=True)
64+
self.assertEqual(len(embeddings), 3)
65+
self.assertEqual(embeddings[0], [0.1, 0.2])
6366

64-
for emb1, emb2 in zip(embeddings1, embeddings2):
65-
self.assertSequenceEqual(emb1, emb2)
67+
def test_exception_handling_aborts_silence(self):
68+
messages = ["Fail 1", "Fail 2"]
69+
70+
with patch.object(self.adapter.model, "encode", side_effect=Exception("API Down")):
71+
with self.assertRaises(Exception) as context:
72+
self.adapter.get_embeddings(messages)
73+
74+
self.assertTrue("API Down" in str(context.exception))
6675

6776
def test_huggingface_model_loaded_once(self):
68-
# Reset class-level cache and counter before test
6977
HuggingFaceAdapter._model_cache.clear()
7078
HuggingFaceAdapter._load_count = 0
7179

72-
# Instantiate multiple adapters
7380
adapter1 = HuggingFaceAdapter(cache=True)
7481
adapter2 = HuggingFaceAdapter(cache=True)
75-
adapter3 = HuggingFaceAdapter(cache=True)
7682

77-
# All adapters should use the same underlying model object
7883
self.assertIs(adapter1.model, adapter2.model)
79-
self.assertIs(adapter2.model, adapter3.model)
80-
81-
# And the model should have been loaded only once
8284
self.assertEqual(HuggingFaceAdapter._load_count, 1)

0 commit comments

Comments
 (0)