From e971dfeaf3e132e59c41767a9883fbdb4f193477 Mon Sep 17 00:00:00 2001 From: Mehmet Can Ay Date: Tue, 9 Dec 2025 10:49:02 +0100 Subject: [PATCH] fix: load HuggingFace models once --- datastew/embedding.py | 10 +++++++++- tests/test_embedding.py | 17 +++++++++++++++++ 2 files changed, 26 insertions(+), 1 deletion(-) diff --git a/datastew/embedding.py b/datastew/embedding.py index 12d929a..7f5c90d 100644 --- a/datastew/embedding.py +++ b/datastew/embedding.py @@ -171,6 +171,9 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]: class HuggingFaceAdapter(EmbeddingModel): + _model_cache = {} + _load_count = 0 # For testing + 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. @@ -178,7 +181,12 @@ def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2", c :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. diff --git a/tests/test_embedding.py b/tests/test_embedding.py index b418a84..7cc29c7 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -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)