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