"""
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]