CoolFace
Apppublic

RohanExploit/Meta-hackathon

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
retriever.py244 linesDownload Raw Back to pipeline
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