Fix reranker config and added huggingface reranker
This commit is contained in:
@@ -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")
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
|
||||
+20
-4
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
@@ -13,18 +15,33 @@ 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]]:
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user