Skip to content

Commit e971dfe

Browse files
committed
fix: load HuggingFace models once
1 parent 579c165 commit e971dfe

2 files changed

Lines changed: 26 additions & 1 deletion

File tree

datastew/embedding.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -171,14 +171,22 @@ def get_embeddings(self, messages: List[str]) -> Sequence[Sequence[float]]:
171171

172172

173173
class HuggingFaceAdapter(EmbeddingModel):
174+
_model_cache = {}
175+
_load_count = 0 # For testing
176+
174177
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2", cache: bool = False):
175178
"""Initialize the Hugging Face adapter with a specified model name.
176179
177180
:param model_name: The model name for sentence transformers, defaults to sentence-transformers/all-MiniLM-L6-v2.
178181
:param cache: Enable or disable caching, defaults to False.
179182
"""
180183
super().__init__(model_name, cache)
181-
self.model = SentenceTransformer(model_name)
184+
185+
if model_name not in self._model_cache:
186+
HuggingFaceAdapter._load_count += 1
187+
self._model_cache[model_name] = SentenceTransformer(model_name)
188+
189+
self.model = self._model_cache[model_name]
182190

183191
def get_embedding(self, text: str) -> Sequence[float]:
184192
"""Retrieve an embedding for a single text input using MPnet.

tests/test_embedding.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -63,3 +63,20 @@ def test_caching_get_embeddings(self):
6363

6464
for emb1, emb2 in zip(embeddings1, embeddings2):
6565
self.assertSequenceEqual(emb1, emb2)
66+
67+
def test_huggingface_model_loaded_once(self):
68+
# Reset class-level cache and counter before test
69+
HuggingFaceAdapter._model_cache.clear()
70+
HuggingFaceAdapter._load_count = 0
71+
72+
# Instantiate multiple adapters
73+
adapter1 = HuggingFaceAdapter(cache=True)
74+
adapter2 = HuggingFaceAdapter(cache=True)
75+
adapter3 = HuggingFaceAdapter(cache=True)
76+
77+
# All adapters should use the same underlying model object
78+
self.assertIs(adapter1.model, adapter2.model)
79+
self.assertIs(adapter2.model, adapter3.model)
80+
81+
# And the model should have been loaded only once
82+
self.assertEqual(HuggingFaceAdapter._load_count, 1)

0 commit comments

Comments
 (0)