diff --git a/mem0/embeddings/configs.py b/mem0/embeddings/configs.py index b4fadd68c..b924263eb 100644 --- a/mem0/embeddings/configs.py +++ b/mem0/embeddings/configs.py @@ -24,6 +24,7 @@ class EmbedderConfig(BaseModel): "lmstudio", "langchain", "aws_bedrock", + "fastembed", ]: return v else: diff --git a/mem0/embeddings/fastembed.py b/mem0/embeddings/fastembed.py new file mode 100644 index 000000000..83868f283 --- /dev/null +++ b/mem0/embeddings/fastembed.py @@ -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] diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 5aa98a158..534a9322b 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 26b7069dc..69415e13a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/tests/embeddings/test_fastembed_embeddings.py b/tests/embeddings/test_fastembed_embeddings.py new file mode 100644 index 000000000..c9d1299bd --- /dev/null +++ b/tests/embeddings/test_fastembed_embeddings.py @@ -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] \ No newline at end of file