[feat add]FastEmbed embedding for local embeddings (#3552)
This commit is contained in:
@@ -24,6 +24,7 @@ class EmbedderConfig(BaseModel):
|
||||
"lmstudio",
|
||||
"langchain",
|
||||
"aws_bedrock",
|
||||
"fastembed",
|
||||
]:
|
||||
return v
|
||||
else:
|
||||
|
||||
@@ -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]
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user