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