from __future__ import annotations import logging from retrieval.index import SearchResult logger = logging.getLogger(__name__) _MODEL_NAME = "cross-encoder/ms-marco-MiniLM-L-6-v2" class Reranker: """Cross-encoder re-ranker that lazy-loads its model on first use.""" def __init__(self) -> None: self._model = None def _load(self): if self._model is None: from sentence_transformers import CrossEncoder # noqa: PLC0415 logger.info("Loading cross-encoder model %s …", _MODEL_NAME) self._model = CrossEncoder(_MODEL_NAME) logger.info("Cross-encoder model loaded.") return self._model def rerank( self, query: str, results: list[SearchResult], top_k: int = 5, ) -> list[SearchResult]: """Re-rank results using cross-encoder scores. Args: query: The user's question string. results: Candidate SearchResults (typically 20 from hybrid search). top_k: How many to return after re-ranking. Returns: Up to top_k SearchResults ordered by descending cross-encoder score. """ if not results: return results model = self._load() pairs = [(query, r.metadata.get("text", "")) for r in results] scores = model.predict(pairs) ranked = sorted( zip(scores, range(len(results)), results), key=lambda x: x[0], reverse=True, ) reranked: list[SearchResult] = [] for new_pos, (ce_score, orig_pos, result) in enumerate(ranked[:top_k]): if orig_pos != new_pos: logger.debug( "chunk_id=%s original_rank=%d → reranked_pos=%d ce_score=%.4f", result.metadata.get("chunk_id", "?"), orig_pos + 1, new_pos + 1, float(ce_score), ) reranked.append(SearchResult(score=float(ce_score), metadata=result.metadata)) return reranked