Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions datastew/embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,16 +51,16 @@ def add_to_cache(self, text: str, embedding: Sequence[float]):
with self._cache_lock:
self._cache[text] = embedding

def get_from_cache(self, text: str) -> Sequence[float]:
def get_from_cache(self, text: str) -> Optional[Sequence[float]]:
"""Retrieve an embedding from the cache.

:param text: Cached input text.
:return: Embedding of the cached input text.
:return: Embedding of the cached input text or `None` if not present.
"""
if self._cache_lock and self._cache is not None:
with self._cache_lock:
return self._cache.get(text, [])
return []
return self._cache.get(text, None)
return None

def get_cached_embeddings(self, messages: List[str]) -> Tuple[List[Sequence[float]], List[int], List[str]]:
"""Retrieve cached embeddings and identify uncached messages.
Expand Down Expand Up @@ -214,7 +214,7 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
"""
sanitized_messages = [self.sanitize(msg) for msg in messages]
if self._cache:
embeddings, uncached_indices, uncached_messages = self.get_cached_embeddings(messages)
embeddings, uncached_indices, uncached_messages = self.get_cached_embeddings(sanitized_messages)

if uncached_messages:
try:
Expand Down Expand Up @@ -268,7 +268,7 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
sanitized_messages = [self.sanitize(msg) for msg in messages]

if self._cache:
embeddings, uncached_indices, uncached_messages = self.get_cached_embeddings(messages)
embeddings, uncached_indices, uncached_messages = self.get_cached_embeddings(sanitized_messages)

if uncached_messages:
try:
Expand Down
23 changes: 2 additions & 21 deletions tests/test_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ def test_hugging_face_adapter_get_embedding(self):
self.assertIsInstance(embedding, Sequence)

def test_hugging_face_adapter_get_embeddings(self):
messages = ["This is message 1.", "This is message 2."]
messages = [f"This is message {i}." for i in range(20)]
embeddings = self.hugging_face_adapter.get_embeddings(messages)
self.assertIsInstance(embeddings, Sequence)
self.assertEqual(len(embeddings), len(messages))
Expand Down Expand Up @@ -47,7 +47,7 @@ def test_caching_get_embedding(self):
self.assertSequenceEqual(embedding1, embedding2)

def test_caching_get_embeddings(self):
messages = ["This is message 1.", "This is message 2."]
messages = [f"This is message {i}." for i in range(20)]
if self.hugging_face_adapter._cache:
self.hugging_face_adapter._cache.clear()

Expand All @@ -63,22 +63,3 @@ def test_caching_get_embeddings(self):

for emb1, emb2 in zip(embeddings1, embeddings2):
self.assertSequenceEqual(emb1, emb2)

def test_cache_vs_no_cache_performance(self):
messages = ["This is message 1.", "This is message 2."]

adapter_with_cache = HuggingFaceAdapter(cache=True)
if adapter_with_cache._cache:
adapter_with_cache._cache.clear()

start_time = time()
adapter_with_cache.get_embeddings(messages)
first_call_time_with_cache = time() - start_time

adapter_without_cache = HuggingFaceAdapter()

start_time = time()
adapter_without_cache.get_embeddings(messages)
first_call_time_without_cache = time() - start_time

self.assertLess(first_call_time_without_cache, first_call_time_with_cache)
Loading