rag-vs / reranker.py
idnameraj's picture
Upload 2 files
636edca verified
Raw
History Blame Contribute Delete
5.35 kB
"""Cross-encoder reranking for RAG retrieval (stage-2 precision)."""
from __future__ import annotations
import logging
import os
from typing import Any
import pandas as pd
from sentence_transformers import CrossEncoder
from conversation import _normalise_chat_turn
logger = logging.getLogger(__name__)
RERANK_ENABLED = os.environ.get("RAG_RERANK_ENABLED", "true").lower() in ("1", "true", "yes")
RERANK_MODEL = os.environ.get(
"RAG_RERANK_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2"
)
RERANK_CANDIDATES = max(1, min(int(os.environ.get("RAG_RERANK_CANDIDATES", "20")), 32))
RERANK_MAX_CHUNK_CHARS = max(128, int(os.environ.get("RAG_RERANK_MAX_CHUNK_CHARS", "512")))
# ms-marco logits are often negative for weak pairs; -9 avoids over-filtering domain docs.
RERANK_MIN_SCORE = float(os.environ.get("RAG_RERANK_MIN_SCORE", "-9.0"))
RERANK_SKIP_DISTANCE = float(os.environ.get("RAG_RERANK_SKIP_DISTANCE", "0.35"))
RERANK_SKIP_MARGIN = float(os.environ.get("RAG_RERANK_SKIP_MARGIN", "0.15"))
_reranker: CrossEncoder | None = None
def load_reranker() -> CrossEncoder | None:
"""Load cross-encoder at startup when reranking is enabled."""
global _reranker
if not RERANK_ENABLED:
logger.info("Reranking disabled (RAG_RERANK_ENABLED=false)")
return None
try:
_reranker = CrossEncoder(RERANK_MODEL)
logger.info("Reranker loaded: %s", RERANK_MODEL)
return _reranker
except Exception as e:
logger.error("Failed to load reranker %s: %s", RERANK_MODEL, e)
_reranker = None
return None
def unload_reranker() -> None:
global _reranker
_reranker = None
def get_reranker() -> CrossEncoder | None:
return _reranker
def reranker_info() -> dict[str, Any]:
return {
"rerank_enabled": RERANK_ENABLED,
"rerank_model": RERANK_MODEL if RERANK_ENABLED else None,
"rerank_loaded": _reranker is not None,
"rerank_candidates": RERANK_CANDIDATES,
"rerank_min_score": RERANK_MIN_SCORE,
"rerank_skip_distance": RERANK_SKIP_DISTANCE,
"rerank_skip_margin": RERANK_SKIP_MARGIN,
}
def _truncate_chunk(text: str) -> str:
text = (text or "").strip()
if len(text) <= RERANK_MAX_CHUNK_CHARS:
return text
return text[: RERANK_MAX_CHUNK_CHARS - 3].rstrip() + "..."
def build_rerank_query(question: str, chat_history: list | None) -> str:
"""
Focused query for the cross-encoder — not the full retrieval blob.
Optionally prefixes the last user turn for short follow-ups.
"""
question = (question or "").strip()
if not question or not chat_history:
return question
last_user = ""
for turn in reversed(chat_history):
pair = _normalise_chat_turn(turn)
if pair and pair[0]:
last_user = pair[0]
break
if last_user and last_user.lower() not in question.lower():
return f"Previous question: {last_user}\nCurrent question: {question}"
return question
def should_skip_rerank(results: pd.DataFrame) -> bool:
"""
Skip cross-encoder when the best vector hit is clearly ahead of the rest.
Saves ~80–150 ms on easy queries.
"""
if not RERANK_ENABLED or get_reranker() is None:
return True
if results.empty or "_distance" not in results.columns:
return True
sorted_df = results.sort_values("_distance", ascending=True)
best = float(sorted_df["_distance"].iloc[0])
if best >= RERANK_SKIP_DISTANCE:
return False
if len(sorted_df) < 2:
return True
second = float(sorted_df["_distance"].iloc[1])
return (second - best) > RERANK_SKIP_MARGIN
def rerank_results(query: str, results: pd.DataFrame, text_col: str = "text") -> pd.DataFrame:
"""Score query–chunk pairs and return rows sorted by rerank score (desc)."""
model = get_reranker()
if model is None or results.empty or not (query or "").strip():
return results
df = results.head(RERANK_CANDIDATES).copy()
texts = [_truncate_chunk(str(row.get(text_col) or "")) for _, row in df.iterrows()]
pairs = [(query, t) for t in texts]
scores = model.predict(pairs, batch_size=32, show_progress_bar=False)
df["_rerank_score"] = scores
df = df.sort_values("_rerank_score", ascending=False)
# Drop weak scores only when some candidates still pass — never return empty here.
filtered = df[df["_rerank_score"] >= RERANK_MIN_SCORE]
if not filtered.empty:
df = filtered
return df
def retrieval_meta_from_results(
results: pd.DataFrame,
*,
path: str,
skipped_rerank: bool,
retrieve_k: int,
) -> dict[str, Any]:
meta: dict[str, Any] = {
"path": path,
"skipped_rerank": skipped_rerank,
"retrieve_k": retrieve_k,
"context_chunks": len(results),
}
if not results.empty and "_distance" in results.columns:
best_row = results.sort_values("_distance", ascending=True).iloc[0]
meta["top_distance"] = float(best_row["_distance"])
if not results.empty and "_rerank_score" in results.columns:
meta["top_rerank_score"] = float(results["_rerank_score"].max())
return meta