From 7e67a38ccbd527f08a9c66afd9deb641b9e08332 Mon Sep 17 00:00:00 2001 From: Mehmet Can Ay Date: Tue, 9 Dec 2025 09:45:48 +0100 Subject: [PATCH 1/2] test(embedding): increase batch size to increase robustness --- tests/test_embedding.py | 23 ++--------------------- 1 file changed, 2 insertions(+), 21 deletions(-) diff --git a/tests/test_embedding.py b/tests/test_embedding.py index d626d92..b418a84 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -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)) @@ -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() @@ -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) From 343d2ba11ff3cc4c5d75311c675cb9a39ab46bc6 Mon Sep 17 00:00:00 2001 From: Mehmet Can Ay Date: Tue, 9 Dec 2025 09:55:04 +0100 Subject: [PATCH 2/2] fix(embedding): return None on cache miss and use sanitized messages --- datastew/embedding.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/datastew/embedding.py b/datastew/embedding.py index 3f411d4..12d929a 100644 --- a/datastew/embedding.py +++ b/datastew/embedding.py @@ -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. @@ -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: @@ -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: