feat: Allow custom model and params with huggingface_base_url (#3574)

This commit is contained in:
Vedant Thakkar
2025-10-14 17:33:18 +05:30
committed by GitHub
parent ce8a285003
commit ea22e8d9cd
2 changed files with 35 additions and 1 deletions
+4 -1
View File
@@ -18,6 +18,7 @@ class HuggingFaceEmbedding(EmbeddingBase):
if config.huggingface_base_url:
self.client = OpenAI(base_url=config.huggingface_base_url)
self.config.model = self.config.model or "tei"
else:
self.config.model = self.config.model or "multi-qa-MiniLM-L6-cos-v1"
@@ -36,6 +37,8 @@ class HuggingFaceEmbedding(EmbeddingBase):
list: The embedding vector.
"""
if self.config.huggingface_base_url:
return self.client.embeddings.create(input=text, model="tei").data[0].embedding
return self.client.embeddings.create(
input=text, model=self.config.model, **self.config.model_kwargs
).data[0].embedding
else:
return self.model.encode(text, convert_to_numpy=True).tolist()
@@ -70,3 +70,34 @@ def test_embed_with_custom_embedding_dims(mock_sentence_transformer):
assert embedder.config.embedding_dims == 768
assert result == [1.0, 1.1, 1.2]
def test_embed_with_huggingface_base_url():
config = BaseEmbedderConfig(
huggingface_base_url="http://localhost:8080",
model="my-custom-model",
model_kwargs={"truncate": True},
)
with patch("mem0.embeddings.huggingface.OpenAI") as mock_openai:
mock_client = Mock()
mock_openai.return_value = mock_client
# Create a mock for the response object and its attributes
mock_embedding_response = Mock()
mock_embedding_response.embedding = [0.1, 0.2, 0.3]
mock_create_response = Mock()
mock_create_response.data = [mock_embedding_response]
mock_client.embeddings.create.return_value = mock_create_response
embedder = HuggingFaceEmbedding(config)
result = embedder.embed("Hello from custom endpoint")
mock_openai.assert_called_once_with(base_url="http://localhost:8080")
mock_client.embeddings.create.assert_called_once_with(
input="Hello from custom endpoint",
model="my-custom-model",
truncate=True,
)
assert result == [0.1, 0.2, 0.3]