From ac9cdd48400ba6303c479dffb6eb55ffc26fcd7c Mon Sep 17 00:00:00 2001 From: Chinnu Abey <152517516+chinnuabey@users.noreply.github.com> Date: Mon, 13 Apr 2026 20:14:22 +0530 Subject: [PATCH] Fix incorrect use of SentenceTransformer for cross-encoder reranker models (#4806) --- mem0/reranker/sentence_transformer_reranker.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/mem0/reranker/sentence_transformer_reranker.py b/mem0/reranker/sentence_transformer_reranker.py index c7a3faf58..2df3b05e6 100644 --- a/mem0/reranker/sentence_transformer_reranker.py +++ b/mem0/reranker/sentence_transformer_reranker.py @@ -6,7 +6,7 @@ from mem0.configs.rerankers.base import BaseRerankerConfig from mem0.configs.rerankers.sentence_transformer import SentenceTransformerRerankerConfig try: - from sentence_transformers import SentenceTransformer + from sentence_transformers import CrossEncoder SENTENCE_TRANSFORMERS_AVAILABLE = True except ImportError: SENTENCE_TRANSFORMERS_AVAILABLE = False @@ -41,7 +41,7 @@ class SentenceTransformerReranker(BaseReranker): ) self.config = config - self.model = SentenceTransformer(self.config.model, device=self.config.device) + self.model = CrossEncoder(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]]: """ @@ -74,8 +74,11 @@ class SentenceTransformerReranker(BaseReranker): # Create query-document pairs pairs = [[query, doc_text] for doc_text in doc_texts] - # Get similarity scores - scores = self.model.predict(pairs) + scores = self.model.predict( + pairs, + batch_size=self.config.batch_size, + show_progress_bar=self.config.show_progress_bar, + ) if isinstance(scores, np.ndarray): scores = scores.tolist()