From 86a2b00c65639e79c86b44d12dc4a39f4ab2e4e4 Mon Sep 17 00:00:00 2001 From: parshvadaftari Date: Tue, 16 Sep 2025 10:48:35 +0530 Subject: [PATCH] Fix reranker config and added huggingface reranker --- mem0/configs/rerankers/base.py | 2 +- mem0/configs/rerankers/cohere.py | 1 - mem0/configs/rerankers/config.py | 4 +- mem0/configs/rerankers/huggingface.py | 17 ++ .../configs/rerankers/sentence_transformer.py | 1 - mem0/memory/main.py | 24 ++- mem0/reranker/huggingface_reranker.py | 147 ++++++++++++++++++ mem0/reranker/llm_reranker.py | 48 ++++-- .../reranker/sentence_transformer_reranker.py | 31 +++- mem0/utils/factory.py | 2 + 10 files changed, 245 insertions(+), 32 deletions(-) create mode 100644 mem0/configs/rerankers/huggingface.py create mode 100644 mem0/reranker/huggingface_reranker.py diff --git a/mem0/configs/rerankers/base.py b/mem0/configs/rerankers/base.py index c48a0a1a0..306dc5b55 100644 --- a/mem0/configs/rerankers/base.py +++ b/mem0/configs/rerankers/base.py @@ -11,7 +11,7 @@ class BaseRerankerConfig(BaseModel): For provider-specific parameters, use the appropriate provider config class. """ - provider: str = Field(description="The reranker provider to use") + provider: Optional[str] = Field(default=None, description="The reranker provider to use") model: Optional[str] = Field(default=None, description="The reranker model to use") api_key: Optional[str] = Field(default=None, description="The API key for the reranker service") top_k: Optional[int] = Field(default=None, description="Maximum number of documents to return after reranking") \ No newline at end of file diff --git a/mem0/configs/rerankers/cohere.py b/mem0/configs/rerankers/cohere.py index d05dc15dc..22105c335 100644 --- a/mem0/configs/rerankers/cohere.py +++ b/mem0/configs/rerankers/cohere.py @@ -10,7 +10,6 @@ class CohereRerankerConfig(BaseRerankerConfig): Inherits from BaseRerankerConfig and adds Cohere-specific settings. """ - provider: str = Field(default="cohere", description="The reranker provider") model: Optional[str] = Field(default="rerank-english-v3.0", description="The Cohere rerank model to use") return_documents: bool = Field(default=False, description="Whether to return the document texts in the response") max_chunks_per_doc: Optional[int] = Field(default=None, description="Maximum number of chunks per document") \ No newline at end of file diff --git a/mem0/configs/rerankers/config.py b/mem0/configs/rerankers/config.py index 46367327d..f0f24aee1 100644 --- a/mem0/configs/rerankers/config.py +++ b/mem0/configs/rerankers/config.py @@ -7,8 +7,8 @@ from mem0.configs.rerankers.base import BaseRerankerConfig class RerankerConfig(BaseModel): """Configuration for rerankers.""" - + provider: str = Field(description="Reranker provider (e.g., 'cohere', 'sentence_transformer')", default="cohere") config: Optional[BaseRerankerConfig] = Field(description="Provider-specific reranker configuration", default=None) - + model_config = {"extra": "forbid"} \ No newline at end of file diff --git a/mem0/configs/rerankers/huggingface.py b/mem0/configs/rerankers/huggingface.py new file mode 100644 index 000000000..1f5794b67 --- /dev/null +++ b/mem0/configs/rerankers/huggingface.py @@ -0,0 +1,17 @@ +from typing import Optional +from pydantic import Field + +from mem0.configs.rerankers.base import BaseRerankerConfig + + +class HuggingFaceRerankerConfig(BaseRerankerConfig): + """ + Configuration class for HuggingFace reranker-specific parameters. + Inherits from BaseRerankerConfig and adds HuggingFace-specific settings. + """ + + model: Optional[str] = Field(default="BAAI/bge-reranker-base", description="The HuggingFace model to use for reranking") + device: Optional[str] = Field(default=None, description="Device to run the model on ('cpu', 'cuda', etc.)") + batch_size: int = Field(default=32, description="Batch size for processing documents") + max_length: int = Field(default=512, description="Maximum length for tokenization") + normalize: bool = Field(default=True, description="Whether to normalize scores") \ No newline at end of file diff --git a/mem0/configs/rerankers/sentence_transformer.py b/mem0/configs/rerankers/sentence_transformer.py index 6339cb74d..70093e6d4 100644 --- a/mem0/configs/rerankers/sentence_transformer.py +++ b/mem0/configs/rerankers/sentence_transformer.py @@ -10,7 +10,6 @@ class SentenceTransformerRerankerConfig(BaseRerankerConfig): Inherits from BaseRerankerConfig and adds Sentence Transformer-specific settings. """ - provider: str = Field(default="sentence_transformer", description="The reranker provider") model: Optional[str] = Field(default="cross-encoder/ms-marco-MiniLM-L-6-v2", description="The cross-encoder model name to use") device: Optional[str] = Field(default=None, description="Device to run the model on ('cpu', 'cuda', etc.)") batch_size: int = Field(default=32, description="Batch size for processing documents") diff --git a/mem0/memory/main.py b/mem0/memory/main.py index a4c3d9dbd..51ab25cd7 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -161,12 +161,28 @@ class Memory(MemoryBase): else: self.graph = None - telemetry_config = deepcopy(self.config.vector_store.config) - telemetry_config.collection_name = "mem0migrations" + # Create telemetry config manually to avoid deepcopy issues with thread locks + telemetry_config_dict = {} + if hasattr(self.config.vector_store.config, 'model_dump'): + # For pydantic models + telemetry_config_dict = self.config.vector_store.config.model_dump() + else: + # For other objects, manually copy common attributes + for attr in ['host', 'port', 'path', 'api_key', 'index_name', 'dimension', 'metric']: + if hasattr(self.config.vector_store.config, attr): + telemetry_config_dict[attr] = getattr(self.config.vector_store.config, attr) + + # Override collection name for telemetry + telemetry_config_dict['collection_name'] = "mem0migrations" + + # Set path for file-based vector stores if self.config.vector_store.provider in ["faiss", "qdrant"]: provider_path = f"migrations_{self.config.vector_store.provider}" - telemetry_config.path = os.path.join(mem0_dir, provider_path) - os.makedirs(telemetry_config.path, exist_ok=True) + telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path) + os.makedirs(telemetry_config_dict['path'], exist_ok=True) + + # Create the config object using the same class as the original + telemetry_config = self.config.vector_store.config.__class__(**telemetry_config_dict) self._telemetry_vector_store = VectorStoreFactory.create( self.config.vector_store.provider, telemetry_config ) diff --git a/mem0/reranker/huggingface_reranker.py b/mem0/reranker/huggingface_reranker.py new file mode 100644 index 000000000..22a3edff4 --- /dev/null +++ b/mem0/reranker/huggingface_reranker.py @@ -0,0 +1,147 @@ +from typing import List, Dict, Any, Optional, Union +import numpy as np + +from mem0.reranker.base import BaseReranker +from mem0.configs.rerankers.base import BaseRerankerConfig +from mem0.configs.rerankers.huggingface import HuggingFaceRerankerConfig + +try: + from transformers import AutoTokenizer, AutoModelForSequenceClassification + import torch + TRANSFORMERS_AVAILABLE = True +except ImportError: + TRANSFORMERS_AVAILABLE = False + + +class HuggingFaceReranker(BaseReranker): + """HuggingFace Transformers based reranker implementation.""" + + def __init__(self, config: Union[BaseRerankerConfig, HuggingFaceRerankerConfig, Dict]): + """ + Initialize HuggingFace reranker. + + Args: + config: Configuration object with reranker parameters + """ + if not TRANSFORMERS_AVAILABLE: + raise ImportError("transformers package is required for HuggingFaceReranker. Install with: pip install transformers torch") + + # Convert to HuggingFaceRerankerConfig if needed + if isinstance(config, dict): + config = HuggingFaceRerankerConfig(**config) + elif isinstance(config, BaseRerankerConfig) and not isinstance(config, HuggingFaceRerankerConfig): + # Convert BaseRerankerConfig to HuggingFaceRerankerConfig with defaults + config = HuggingFaceRerankerConfig( + provider=getattr(config, 'provider', 'huggingface'), + model=getattr(config, 'model', 'BAAI/bge-reranker-base'), + api_key=getattr(config, 'api_key', None), + top_k=getattr(config, 'top_k', None), + device=None, # Will auto-detect + batch_size=32, # Default + max_length=512, # Default + normalize=True, # Default + ) + + self.config = config + + # Set device + if self.config.device is None: + self.device = "cuda" if torch.cuda.is_available() else "cpu" + else: + self.device = self.config.device + + # Load model and tokenizer + self.tokenizer = AutoTokenizer.from_pretrained(self.config.model) + self.model = AutoModelForSequenceClassification.from_pretrained(self.config.model) + self.model.to(self.device) + self.model.eval() + + def rerank(self, query: str, documents: List[Dict[str, Any]], top_k: int = None) -> List[Dict[str, Any]]: + """ + Rerank documents using HuggingFace cross-encoder model. + + Args: + query: The search query + documents: List of documents to rerank + top_k: Number of top documents to return + + Returns: + List of reranked documents with rerank_score + """ + if not documents: + return documents + + # Extract text content for reranking + doc_texts = [] + for doc in documents: + if 'memory' in doc: + doc_texts.append(doc['memory']) + elif 'text' in doc: + doc_texts.append(doc['text']) + elif 'content' in doc: + doc_texts.append(doc['content']) + else: + doc_texts.append(str(doc)) + + try: + scores = [] + + # Process documents in batches + for i in range(0, len(doc_texts), self.config.batch_size): + batch_docs = doc_texts[i:i + self.config.batch_size] + batch_pairs = [[query, doc] for doc in batch_docs] + + # Tokenize batch + inputs = self.tokenizer( + batch_pairs, + padding=True, + truncation=True, + max_length=self.config.max_length, + return_tensors="pt" + ).to(self.device) + + # Get scores + with torch.no_grad(): + outputs = self.model(**inputs) + batch_scores = outputs.logits.squeeze(-1).cpu().numpy() + + # Handle single item case + if batch_scores.ndim == 0: + batch_scores = [float(batch_scores)] + else: + batch_scores = batch_scores.tolist() + + scores.extend(batch_scores) + + # Normalize scores if requested + if self.config.normalize: + scores = np.array(scores) + scores = (scores - scores.min()) / (scores.max() - scores.min() + 1e-8) + scores = scores.tolist() + + # Combine documents with scores + doc_score_pairs = list(zip(documents, scores)) + + # Sort by score (descending) + doc_score_pairs.sort(key=lambda x: x[1], reverse=True) + + # Apply top_k limit + final_top_k = top_k or self.config.top_k + if final_top_k: + doc_score_pairs = doc_score_pairs[:final_top_k] + + # Create reranked results + reranked_docs = [] + for doc, score in doc_score_pairs: + reranked_doc = doc.copy() + reranked_doc['rerank_score'] = float(score) + reranked_docs.append(reranked_doc) + + return reranked_docs + + except Exception as e: + # Fallback to original order if reranking fails + for doc in documents: + doc['rerank_score'] = 0.0 + final_top_k = top_k or self.config.top_k + return documents[:final_top_k] if final_top_k else documents \ No newline at end of file diff --git a/mem0/reranker/llm_reranker.py b/mem0/reranker/llm_reranker.py index 6eff0dbe8..dff55a20a 100644 --- a/mem0/reranker/llm_reranker.py +++ b/mem0/reranker/llm_reranker.py @@ -1,39 +1,55 @@ import os import re -from typing import List, Dict, Any, Optional +from typing import List, Dict, Any, Optional, Union from mem0.reranker.base import BaseReranker from mem0.utils.factory import LlmFactory +from mem0.configs.rerankers.base import BaseRerankerConfig +from mem0.configs.rerankers.llm import LLMRerankerConfig class LLMReranker(BaseReranker): """LLM-based reranker implementation.""" - - def __init__(self, config): + + def __init__(self, config: Union[BaseRerankerConfig, LLMRerankerConfig, Dict]): """ Initialize LLM reranker. - + Args: - config: LLMRerankerConfig object with configuration parameters + config: Configuration object with reranker parameters """ + # Convert to LLMRerankerConfig if needed + if isinstance(config, dict): + config = LLMRerankerConfig(**config) + elif isinstance(config, BaseRerankerConfig) and not isinstance(config, LLMRerankerConfig): + # Convert BaseRerankerConfig to LLMRerankerConfig with defaults + config = LLMRerankerConfig( + provider=getattr(config, 'provider', 'openai'), + model=getattr(config, 'model', 'gpt-4o-mini'), + api_key=getattr(config, 'api_key', None), + top_k=getattr(config, 'top_k', None), + temperature=0.0, # Default for reranking + max_tokens=100, # Default for reranking + ) + self.config = config - + # Create LLM configuration for the factory llm_config = { - "model": config.model, - "temperature": config.temperature, - "max_tokens": config.max_tokens, + "model": self.config.model, + "temperature": self.config.temperature, + "max_tokens": self.config.max_tokens, } - + # Add API key if provided - if config.api_key: - llm_config["api_key"] = config.api_key - + if self.config.api_key: + llm_config["api_key"] = self.config.api_key + # Initialize LLM using the factory - self.llm = LlmFactory.create(config.provider, llm_config) - + self.llm = LlmFactory.create(self.config.provider, llm_config) + # Default scoring prompt - self.scoring_prompt = config.scoring_prompt or self._get_default_prompt() + self.scoring_prompt = getattr(self.config, 'scoring_prompt', None) or self._get_default_prompt() def _get_default_prompt(self) -> str: """Get the default scoring prompt template.""" diff --git a/mem0/reranker/sentence_transformer_reranker.py b/mem0/reranker/sentence_transformer_reranker.py index efe9ca8fa..f849d66e2 100644 --- a/mem0/reranker/sentence_transformer_reranker.py +++ b/mem0/reranker/sentence_transformer_reranker.py @@ -1,7 +1,9 @@ -from typing import List, Dict, Any, Optional +from typing import List, Dict, Any, Optional, Union import numpy as np from mem0.reranker.base import BaseReranker +from mem0.configs.rerankers.base import BaseRerankerConfig +from mem0.configs.rerankers.sentence_transformer import SentenceTransformerRerankerConfig try: from sentence_transformers import SentenceTransformer, util @@ -12,19 +14,34 @@ except ImportError: class SentenceTransformerReranker(BaseReranker): """Sentence Transformer based reranker implementation.""" - - def __init__(self, config): + + def __init__(self, config: Union[BaseRerankerConfig, SentenceTransformerRerankerConfig, Dict]): """ Initialize Sentence Transformer reranker. - + Args: - config: SentenceTransformerRerankerConfig object with configuration parameters + config: Configuration object with reranker parameters """ if not SENTENCE_TRANSFORMERS_AVAILABLE: raise ImportError("sentence-transformers package is required for SentenceTransformerReranker. Install with: pip install sentence-transformers") - + + # Convert to SentenceTransformerRerankerConfig if needed + if isinstance(config, dict): + config = SentenceTransformerRerankerConfig(**config) + elif isinstance(config, BaseRerankerConfig) and not isinstance(config, SentenceTransformerRerankerConfig): + # Convert BaseRerankerConfig to SentenceTransformerRerankerConfig with defaults + config = SentenceTransformerRerankerConfig( + provider=getattr(config, 'provider', 'sentence_transformer'), + model=getattr(config, 'model', 'cross-encoder/ms-marco-MiniLM-L-6-v2'), + api_key=getattr(config, 'api_key', None), + top_k=getattr(config, 'top_k', None), + device=None, # Will auto-detect + batch_size=32, # Default + show_progress_bar=False, # Default + ) + self.config = config - self.model = SentenceTransformer(config.model, device=config.device) + self.model = SentenceTransformer(self.config.model, device=self.config.device) def rerank(self, query: str, documents: List[Dict[str, Any]], top_k: int = None) -> List[Dict[str, Any]]: """ diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index b279ed979..ff24956ce 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -15,6 +15,7 @@ from mem0.configs.rerankers.cohere import CohereRerankerConfig from mem0.configs.rerankers.sentence_transformer import SentenceTransformerRerankerConfig from mem0.configs.rerankers.zero_entropy import ZeroEntropyRerankerConfig from mem0.configs.rerankers.llm import LLMRerankerConfig +from mem0.configs.rerankers.huggingface import HuggingFaceRerankerConfig from mem0.embeddings.mock import MockEmbeddings @@ -235,6 +236,7 @@ class RerankerFactory: "sentence_transformer": ("mem0.reranker.sentence_transformer_reranker.SentenceTransformerReranker", SentenceTransformerRerankerConfig), "zero_entropy": ("mem0.reranker.zero_entropy_reranker.ZeroEntropyReranker", ZeroEntropyRerankerConfig), "llm": ("mem0.reranker.llm_reranker.LLMReranker", LLMRerankerConfig), + "huggingface": ("mem0.reranker.huggingface_reranker.HuggingFaceReranker", HuggingFaceRerankerConfig), } @classmethod