rag-document-qa / retrieval /reranker.py
Amrita P
feat: implement advanced RAG pipeline (cross-encoder, contextual chunks, streaming, confidence gating)
4f25e4a
Raw
History Blame Contribute Delete
2.1 kB
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