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
10 changes: 9 additions & 1 deletion datastew/embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,14 +171,22 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:


class HuggingFaceAdapter(EmbeddingModel):
_model_cache = {}
_load_count = 0 # For testing
Comment thread
tiadams marked this conversation as resolved.

def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2", cache: bool = False):
"""Initialize the Hugging Face adapter with a specified model name.

:param model_name: The model name for sentence transformers, defaults to sentence-transformers/all-MiniLM-L6-v2.
:param cache: Enable or disable caching, defaults to False.
"""
super().__init__(model_name, cache)
self.model = SentenceTransformer(model_name)

if model_name not in self._model_cache:
HuggingFaceAdapter._load_count += 1
self._model_cache[model_name] = SentenceTransformer(model_name)

self.model = self._model_cache[model_name]

def get_embedding(self, text: str) -> Sequence[float]:
"""Retrieve an embedding for a single text input using MPnet.
Expand Down
17 changes: 17 additions & 0 deletions tests/test_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,3 +63,20 @@ def test_caching_get_embeddings(self):

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

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

# Instantiate multiple adapters
adapter1 = HuggingFaceAdapter(cache=True)
adapter2 = HuggingFaceAdapter(cache=True)
adapter3 = HuggingFaceAdapter(cache=True)

# All adapters should use the same underlying model object
self.assertIs(adapter1.model, adapter2.model)
self.assertIs(adapter2.model, adapter3.model)

# And the model should have been loaded only once
self.assertEqual(HuggingFaceAdapter._load_count, 1)
Loading