RohanExploit/Meta-hackathon
0
1"""Two-stage retriever — simulating NVIDIA NeMo Reranker on free-tier infra.2 3Stage 1 (Dense Retrieval):4 FAISS similarity_search with HuggingFaceEmbeddings → top-20 candidates.5 6Stage 2 (Cross-Encoder Reranking via LLM):7 Groq's llama-3.3-70b-versatile scores each of the 20 chunks against the8 query on a 0-10 relevance scale, then we keep the top 3-5 by score.9 10This two-stage approach achieves near cross-encoder recall quality without11requiring a GPU-resident reranker model.12"""13 14from __future__ import annotations15 16import asyncio17import json18import logging19from dataclasses import dataclass20from typing import List, Optional21 22from langchain_community.vectorstores import FAISS23from langchain_core.documents import Document24from langchain_huggingface import HuggingFaceEmbeddings25from groq import AsyncGroq26 27from .config import get_config28 29logger = logging.getLogger(__name__)30 31 32@dataclass33class RankedChunk:34 """A document chunk with its relevance score after reranking."""35 36 content: str37 score: float38 metadata: dict39 40 41# ── Embedding model singleton ────────────────────────────────────────42 43_embeddings: HuggingFaceEmbeddings | None = None44 45 46def get_embeddings() -> HuggingFaceEmbeddings:47 """Return the shared HuggingFace embedding model."""48 global _embeddings49 if _embeddings is None:50 cfg = get_config()51 _embeddings = HuggingFaceEmbeddings(52 model_name=cfg.embedding.model_name,53 model_kwargs={"device": cfg.embedding.device},54 encode_kwargs={"normalize_embeddings": True},55 )56 return _embeddings57 58 59# ── FAISS vector store ───────────────────────────────────────────────60 61_vectorstore: FAISS | None = None62 63 64def load_vectorstore(path: str | None = None) -> FAISS:65 """Load FAISS index from disk. Creates singleton."""66 global _vectorstore67 if _vectorstore is None:68 cfg = get_config()69 index_path = path or cfg.retriever.faiss_index_path70 _vectorstore = FAISS.load_local(71 index_path,72 get_embeddings(),73 allow_dangerous_deserialization=True,74 )75 logger.info("Loaded FAISS index from %s", index_path)76 return _vectorstore77 78 79def set_vectorstore(vs: FAISS) -> None:80 """Inject a pre-built vectorstore (used during ingestion or testing)."""81 global _vectorstore82 _vectorstore = vs83 84 85# ── Stage 1: Dense Retrieval ─────────────────────────────────────────86 87async def dense_retrieve(query: str, k: int | None = None) -> List[Document]:88 """Fetch top-k documents from FAISS using dense embedding similarity.89 90 Runs the CPU-bound FAISS search in a thread pool to keep the event91 loop unblocked.92 """93 cfg = get_config()94 top_k = k or cfg.retriever.dense_top_k95 vs = load_vectorstore()96 97 # FAISS search is CPU-bound — offload to a thread98 loop = asyncio.get_event_loop()99 docs = await loop.run_in_executor(100 None,101 lambda: vs.similarity_search(query, k=top_k),102 )103 logger.info("Dense retrieval returned %d documents for query", len(docs))104 return docs105 106 107# ── Stage 2: LLM Cross-Encoder Reranking ─────────────────────────────108 109RERANK_PROMPT = """\110You are a relevance scoring engine. Given a user query and a text chunk, \111rate the chunk's relevance to the query on a scale of 0 to 10.112 113Rules:114- 0 means completely irrelevant.115- 10 means perfectly answers the query.116- Consider semantic relevance, not just keyword overlap.117- Respond with ONLY a JSON object: {{"score": <number>, "reason": "<brief reason>"}}118 119User Query: {query}120 121Text Chunk:122---123{chunk}124---125 126Your relevance score:"""127 128 129async def _score_chunk(130 client: AsyncGroq,131 model: str,132 query: str,133 doc: Document,134 semaphore: asyncio.Semaphore,135) -> RankedChunk:136 """Score a single chunk using the LLM reranker."""137 async with semaphore:138 prompt = RERANK_PROMPT.format(query=query, chunk=doc.page_content[:1500])139 try:140 response = await client.chat.completions.create(141 model=model,142 messages=[{"role": "user", "content": prompt}],143 temperature=0.0,144 max_tokens=80,145 )146 raw = (response.choices[0].message.content or "").strip()147 148 # Parse score from JSON149 try:150 data = json.loads(raw)151 score = float(data.get("score", 0))152 except (json.JSONDecodeError, TypeError, ValueError):153 # Fallback: extract first number154 import re155 match = re.search(r"(\d+(?:\.\d+)?)", raw)156 score = float(match.group(1)) if match else 0.0157 158 # Normalize to 0-1159 score = max(0.0, min(10.0, score)) / 10.0160 161 except Exception as exc:162 logger.warning("Rerank scoring failed for chunk: %s", exc)163 score = 0.0164 165 return RankedChunk(166 content=doc.page_content,167 score=score,168 metadata=doc.metadata or {},169 )170 171 172async def rerank(173 query: str,174 documents: List[Document],175 top_k: int | None = None,176) -> List[RankedChunk]:177 """Score and rerank documents using LLM cross-encoder pattern.178 179 Fires all scoring requests concurrently (bounded by semaphore) for180 minimum wall-clock latency.181 182 Args:183 query: The user's search query.184 documents: Candidate documents from Stage 1.185 top_k: Number of top results to return (default from config).186 187 Returns:188 Sorted list of RankedChunk, highest relevance first.189 """190 cfg = get_config()191 final_k = top_k or cfg.retriever.rerank_top_k192 client = AsyncGroq(api_key=cfg.groq.api_key)193 194 # Limit concurrent Groq API calls to avoid rate limiting195 semaphore = asyncio.Semaphore(cfg.max_concurrent_llm_calls)196 197 # Score all chunks concurrently198 tasks = [199 _score_chunk(client, cfg.groq.reranker_model, query, doc, semaphore)200 for doc in documents201 ]202 ranked = await asyncio.gather(*tasks)203 204 # Sort by score descending, filter by minimum threshold, take top_k205 ranked.sort(key=lambda r: r.score, reverse=True)206 filtered = [r for r in ranked if r.score >= cfg.retriever.rerank_min_score]207 208 result = filtered[:final_k]209 logger.info(210 "Reranking: %d candidates → %d after threshold → returning top %d",211 len(documents),212 len(filtered),213 len(result),214 )215 return result216 217 218# ── Combined Two-Stage Pipeline ──────────────────────────────────────219 220async def retrieve_and_rerank(221 query: str,222 dense_k: int | None = None,223 final_k: int | None = None,224) -> List[RankedChunk]:225 """Full two-stage retrieval: dense fetch → LLM rerank.226 227 Args:228 query: User's search query.229 dense_k: Number of candidates for Stage 1 (default: 20).230 final_k: Number of results after reranking (default: 5).231 232 Returns:233 Top-k RankedChunks sorted by relevance.234 """235 # Stage 1: Dense retrieval236 candidates = await dense_retrieve(query, k=dense_k)237 238 if not candidates:239 logger.warning("No documents found in dense retrieval")240 return []241 242 # Stage 2: LLM reranking243 return await rerank(query, candidates, top_k=final_k)244 