import os import logging from typing import List, Dict import numpy as np from sentence_transformers import CrossEncoder from groq import Groq from vector_store import ( semantic_search, bm25_search, hybrid_search, ) # Logger settings logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) console_handler = logging.StreamHandler() formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s') console_handler.setFormatter(formatter) logger.addHandler(console_handler) # Initialization def init_reranker(model_name: str = "cross-encoder/ms-marco-MiniLM-L-6-v2") -> CrossEncoder: """Initialize CrossEncoder for document reordering.""" logger.info(f"CrossEncoder initialization: {model_name}") return CrossEncoder(model_name) def init_groq(api_key: str = None) -> Groq: """Initializing the GROQ client for LLM requests.""" if api_key: os.environ["GROQ_API_KEY"] = api_key client = Groq(api_key=os.environ.get("GROQ_API_KEY")) logger.info("GROQ client initialized") return client # Reranking def rerank(query: str, docs: List[Dict], reranker: CrossEncoder, top_k: int = 5) -> List[Dict]: """Reranking documents using CrossEncoder based on a query.""" if not docs: return [] pairs = [[query, d["text"]] for d in docs] scores = reranker.predict(pairs) for i, s in enumerate(scores): docs[i]["rerank_score"] = float(s) ranked = sorted(docs, key=lambda x: x["rerank_score"], reverse=True) return ranked[:top_k] # LLM answering def llm_answer(query: str, context: List[Dict], client: Groq) -> str: """Forming an LLM response based on the provided document context.""" context_text = "\n\n---\n\n".join(f"[{d['id']}] {d['text']}" for d in context) prompt = f""" You are an AI assistant answering based only on provided context. Question: {query} Context: {context_text} Answer only using information from the context. If answer not found, say "I don't know". """ completion = client.chat.completions.create( model="llama-3.3-70b-versatile", messages=[{"role": "user", "content": prompt}], temperature=0 ) return completion.choices[0].message.content # Retrieve documents def retrieve_documents( query: str, documents: list, model, index, bm25, mode: str = "semantic", k: int = 20 ) -> List[Dict]: """Retrieve documents using semantic, bm25 or hybrid mode.""" if mode == "semantic": raw_results = semantic_search(query, model, index, documents, k=k) elif mode == "bm25": raw_results = bm25_search(query, bm25, documents, k=k) elif mode == "hybrid": raw_results = hybrid_search(query, model, index, bm25, documents, k=k) else: raise ValueError(f"Unknown retrieve mode: {mode}") docs = [] for item in raw_results: doc = item["document"] docs.append({ "id": doc["id"], "text": doc["text"], "metadata": doc.get("metadata", {}), "score": item["score"] }) return docs # Full RAG Pipeline def rag_pipeline( query: str, reranker_model: CrossEncoder, llm_client: Groq, documents: list, model, faiss_index, bm25_index, retrieve_mode: str = "hybrid", retrieve_k: int = 20, rerank_k: int = 5 ) -> Dict: """ Full RAG pipeline: document retrieval, reranking and LLM response generation. """ # 1. Retrieve initial_docs = retrieve_documents( query, documents=documents, model=model, index=faiss_index, bm25=bm25_index, mode=retrieve_mode, k=retrieve_k ) logger.info(f"Documents after retrieve: {len(initial_docs)}") # 2. Rerank reranked_docs = rerank(query, initial_docs, reranker_model, top_k=rerank_k) logger.info(f"Documents after rerank: {len(reranked_docs)}") # 3. LLM final answer answer = llm_answer(query, reranked_docs, llm_client) logger.info("Received response from LLM") return { "query": query, "retrieved_docs": initial_docs, "reranked_docs": reranked_docs, "answer": answer }