[feat add]FastEmbed embedding for local embeddings (#3552)

This commit is contained in:
Tarun Jain
2025-10-16 22:52:22 +05:30
committed by GitHub
parent 394203d1b5
commit 1090784302
5 changed files with 78 additions and 0 deletions
+1
View File
@@ -24,6 +24,7 @@ class EmbedderConfig(BaseModel):
"lmstudio",
"langchain",
"aws_bedrock",
"fastembed",
]:
return v
else:
+29
View File
@@ -0,0 +1,29 @@
from typing import Optional, Literal
from mem0.embeddings.base import EmbeddingBase
from mem0.configs.embeddings.base import BaseEmbedderConfig
try:
from fastembed import TextEmbedding
except ImportError:
raise ImportError("FastEmbed is not installed. Please install it using `pip install fastembed`")
class FastEmbedEmbedding(EmbeddingBase):
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
super().__init__(config)
self.config.model = self.config.model or "thenlper/gte-large"
self.dense_model = TextEmbedding(model_name = self.config.model)
def embed(self, text, memory_action: Optional[Literal["add", "search", "update"]] = None):
"""
Convert the text to embeddings using FastEmbed running in the Onnx runtime
Args:
text (str): The text to embed.
memory_action (optional): The type of embedding to use. Must be one of "add", "search", or "update". Defaults to None.
Returns:
list: The embedding vector.
"""
text = text.replace("\n", " ")
embeddings = list(self.dense_model.embed(text))
return embeddings[0]
+1
View File
@@ -145,6 +145,7 @@ class EmbedderFactory:
"lmstudio": "mem0.embeddings.lmstudio.LMStudioEmbedding",
"langchain": "mem0.embeddings.langchain.LangchainEmbedding",
"aws_bedrock": "mem0.embeddings.aws_bedrock.AWSBedrockEmbedding",
"fastembed": "mem0.embeddings.fastembed.FastEmbedEmbedding",
}
@classmethod
+1
View File
@@ -72,6 +72,7 @@ extras = [
"sentence-transformers>=5.0.0",
"elasticsearch>=8.0.0,<9.0.0",
"opensearch-py>=2.0.0",
"fastembed>=0.3.1",
]
test = [
"pytest>=8.2.2",
@@ -0,0 +1,46 @@
from unittest.mock import Mock, patch
import pytest
import numpy as np
from mem0.configs.embeddings.base import BaseEmbedderConfig
try:
from mem0.embeddings.fastembed import FastEmbedEmbedding
except ImportError:
pytest.skip("fastembed not installed", allow_module_level=True)
@pytest.fixture
def mock_fastembed_client():
with patch("mem0.embeddings.fastembed.TextEmbedding") as mock_fastembed:
mock_client = Mock()
mock_fastembed.return_value = mock_client
yield mock_client
def test_embed_with_jina_model(mock_fastembed_client):
config = BaseEmbedderConfig(model="jinaai/jina-embeddings-v2-base-en", embedding_dims=768)
embedder = FastEmbedEmbedding(config)
mock_embedding = np.array([0.1, 0.2, 0.3, 0.4, 0.5])
mock_fastembed_client.embed.return_value = iter([mock_embedding])
text = "Sample text to embed."
embedding = embedder.embed(text)
mock_fastembed_client.embed.assert_called_once_with(text)
assert list(embedding) == [0.1, 0.2, 0.3, 0.4, 0.5]
def test_embed_removes_newlines(mock_fastembed_client):
config = BaseEmbedderConfig(model="jinaai/jina-embeddings-v2-base-en", embedding_dims=768)
embedder = FastEmbedEmbedding(config)
mock_embedding = np.array([0.7, 0.8, 0.9])
mock_fastembed_client.embed.return_value = iter([mock_embedding])
text_with_newlines = "Hello\nworld"
embedding = embedder.embed(text_with_newlines)
mock_fastembed_client.embed.assert_called_once_with("Hello world")
assert list(embedding) == [0.7, 0.8, 0.9]