""" Cross-Encoder Reranker Uses ms-marco-MiniLM-L6-v2 to rerank retrieved passages by scoring each (query, passage) pair with a cross-encoder. """ from __future__ import annotations import logging from typing import Dict, List, Optional logger = logging.getLogger(__name__) _cross_encoder = None def _get_model(model_name: str = None): """Load the cross-encoder model""" global _cross_encoder if _cross_encoder is not None: return _cross_encoder try: from . import config as _cfg except ImportError: import config as _cfg model_name = model_name or _cfg.RERANKER_MODEL from sentence_transformers import CrossEncoder logger.info(f"Loading cross-encoder model: {model_name}") _cross_encoder = CrossEncoder(model_name) logger.info("Cross-encoder model loaded.") return _cross_encoder def rerank( query: str, chunks: List[Dict], top_n: Optional[int] = None, model_name: Optional[str] = None, ) -> List[Dict]: """ rerank chunks using the cross-encoder. args: query: The user query string. chunks: List of chunk dicts, each must have a 'text' field. top_n: Number of top passages to return. If None, uses config.RERANKER_TOP_N. model_name: Override model name (uses config default if None). returns: top N chunks sorted by cross-encoder score, each with 'cross_encoder_score' added. """ if not chunks: return [] try: from . import config as _cfg except ImportError: import config as _cfg top_n = top_n if top_n is not None else _cfg.RERANKER_TOP_N model = _get_model(model_name) # build (query, passage) pairs pairs = [(query, chunk.get("text", "")) for chunk in chunks] # score all pairs scores = model.predict(pairs) # attach scores and sort descending scored_chunks = [] for chunk, score in zip(chunks, scores): enriched = chunk.copy() enriched["cross_encoder_score"] = round(float(score), 6) scored_chunks.append(enriched) scored_chunks.sort(key=lambda x: x["cross_encoder_score"], reverse=True) return scored_chunks[:top_n]