RAG-Pipeline / app.py
Sayandip's picture
Update app.py
99bcfcc verified
Raw
History Blame Contribute Delete
10.7 kB
import os
import chromadb
import gradio as gr
from pathlib import Path
from typing import List, Dict, Any
from dataclasses import dataclass
from sentence_transformers import SentenceTransformer
from google import genai
from google.genai import types
# --- 1. CONFIGURATION AND INITIALIZATION ---
API_KEY = os.getenv("GEMINI_API_KEY")
if not API_KEY:
raise ValueError("Missing GEMINI_API_KEY in environment variables or Hugging Face secrets.")
client = genai.Client(api_key=API_KEY)
generation_config = types.GenerateContentConfig(
temperature=0.1,
)
# --- 2. RAG CORE CLASSES (Final Corrected Versions) ---
@dataclass
class Document:
text: str
metadata: Dict[str, Any]
doc_id: str
class VectorStore:
def __init__(self, collection_name: str = "rag_collection", model_name: str = "all-MiniLM-L6-v2"):
self.persist_dir = "./chroma_db"
Path(self.persist_dir).mkdir(exist_ok=True)
# Load the Sentence Transformer model
self.embedding_model = SentenceTransformer(model_name)
# Initialize Chroma Client
self.client = chromadb.PersistentClient(path=self.persist_dir)
self.collection = self.client.get_or_create_collection(
name=collection_name,
embedding_function=self._create_embedding_function()
)
print("VectorStore initialized and ready.")
def _create_embedding_function(self):
# Inner class to wrap the Sentence Transformer model for ChromaDB
class SentenceTransformerEmbeddingFunction(chromadb.api.types.EmbeddingFunction):
def __init__(self, model):
self.embedding_model = model
def __call__(self, texts: List[str]) -> List[List[float]]:
# This function is called by ChromaDB to generate embeddings
return self.embedding_model.encode(texts, convert_to_tensor=False).tolist()
# Pass the model instance to the inner class
return SentenceTransformerEmbeddingFunction(self.embedding_model)
def add_documents(self, documents: List[Document]):
texts = [doc.text for doc in documents]
metadatas = [doc.metadata for doc in documents]
ids = [doc.doc_id for doc in documents]
self.collection.upsert(
documents=texts,
metadatas=metadatas,
ids=ids
)
print(f"Added/Updated **{len(documents)}** documents in the vector store.")
def search(self, query: str, k: int = 4) -> List[Document]:
try:
# FIX: Removed 'ids' from the include parameter. 'ids' is returned by default.
results = self.collection.query(
query_texts=[query],
n_results=k,
include=['documents', 'metadatas']
)
except Exception as e:
print(f"ChromaDB search failed: {e}")
return []
retrieved_docs = []
# Ensure the required keys exist and have data
if results and 'documents' in results and results['documents'] and results['ids']:
# results['ids'][0] contains the IDs corresponding to the retrieved documents
for text, metadata, doc_id in zip(results['documents'][0], results['metadatas'][0], results['ids'][0]):
retrieved_docs.append(Document(text=text, metadata=metadata, doc_id=doc_id))
return retrieved_docs
class RAGPipeline:
def __init__(self, client: genai.Client, vector_store: VectorStore, config: types.GenerateContentConfig):
self.client = client
self.vector_store = vector_store
self.config = config
self.model = "gemini-2.5-flash"
self.system_prompt = (
"You are an expert Question Answering system. "
"Your task is to answer the user's question only based on the provided CONTEXT. "
"If the CONTEXT does not contain the answer, state that you cannot find the answer in the provided documents. "
"Always cite the source document ID (e.g., [chunk1]) after each piece of information you use. "
"Format your response using Markdown."
)
def _create_prompt(self, query: str, context: List[Document]) -> str:
context_text = "\n---\n".join([
f"Document ID: {doc.doc_id} (Source: {doc.metadata.get('source', 'Unknown')})\nContent: {doc.text}"
for doc in context
])
prompt = (
f"{self.system_prompt}\n\n"
f"CONTEXT:\n{context_text}\n\n"
f"USER QUESTION: {query}"
)
return prompt
def generate_response(self, query: str, k: int = 4) -> str:
# --- FIX APPLIED: Removed the check for self.client.api_key which caused the AttributeError ---
# The key is checked during the initial Client setup, and the try/except block
# below will handle runtime errors if the key is invalid or permissions are missing.
# 1. Retrieval
retrieved_docs = self.vector_store.search(query, k=k)
if not retrieved_docs:
return "Retrieval failed. No relevant documents found to answer the query."
# 2. Augmentation & Prompt Creation
final_prompt = self._create_prompt(query, retrieved_docs)
# 3. Generation
try:
# Use keyword argument 'text=final_prompt'
contents = [types.Content(parts=[types.Part.from_text(text=final_prompt)], role="user")]
response = self.client.models.generate_content(
model=self.model,
contents=contents,
config=self.config
)
# Display the context used below the final answer for transparency
context_display = "\n\n---\n\n"
context_display += "**Context Retrieved:**\n"
for doc in retrieved_docs:
context_display += f"* **ID {doc.doc_id}** (Source: {doc.metadata.get('source', 'Unknown')}): {doc.text[:80]}...\n"
return response.text + context_display
except Exception as e:
# This catch-all handles errors from the API call itself (e.g., invalid API key, quota issues)
return f"**Gemini Generation Error:** The API call failed. Ensure your GEMINI_API_KEY is valid and has not expired. Details: {e}"
# --- 3. DATA INDEXING (Runs on app startup) ---
def initialize_rag_system() -> RAGPipeline:
"""Initializes the VectorStore and indexes the data."""
print("--- Initializing RAG System ---")
# Define Sample Knowledge Base
knowledge_base = {
"doc1": "The **Gemini 2.5 Pro** model is highly capable for complex reasoning tasks and can handle up to 2 million tokens of context, making it suitable for long-form content analysis. It was released in 2024.",
"doc2": "Google's **Gemini 2.5 Flash** model is optimized for speed and efficiency, making it ideal for high-volume, low-latency tasks like summarization and chat. It also has a 2 million token context window.",
"doc3": "The **Retrieval-Augmented Generation (RAG)** architecture improves LLM responses by fetching relevant, external documents and providing them as context to the model, reducing hallucinations and improving grounding.",
"doc4": "The new **google-genai** SDK provides a unified interface for interacting with the entire Gemini family of models, including features like function calling and streaming. This SDK is the recommended way to access the latest features.",
"doc5": "ChromaDB is used in this pipeline as the **vector store** to efficiently manage and search the document embeddings. It leverages **Sentence Transformers** for local embedding generation."
}
# Chunking and Document Creation
rag_documents: List[Document] = []
chunk_id_counter = 0
for doc_id, text in knowledge_base.items():
chunk_id_counter += 1
rag_documents.append(
Document(
text=text,
metadata={"source": doc_id, "chunk_index": chunk_id_counter},
doc_id=f"chunk{chunk_id_counter}"
)
)
print(f"Created **{len(rag_documents)}** individual documents/chunks.")
# Initialize Vector Store and Index
vector_store = VectorStore(collection_name="rag_knowledge_base")
vector_store.add_documents(rag_documents)
print("--- Indexing Complete ---")
# Initialize RAG Pipeline
rag_pipeline = RAGPipeline(
client=client,
vector_store=vector_store,
config=generation_config
)
return rag_pipeline
# Initialize the RAG system globally when the app starts
try:
rag_system = initialize_rag_system()
except Exception as e:
# This might catch errors during model loading if no network is available, etc.
print(f"FATAL ERROR during RAG system initialization: {e}")
rag_system = None
# --- 4. GRADIO INTERFACE ---
def query_rag(query: str, k_chunks: int) -> str:
"""Interface function to run the RAG pipeline."""
if rag_system is None:
return "**FATAL ERROR:** The RAG system failed to initialize. Check the console logs for initialization errors (e.g., API key, model downloads)."
if not query:
return "Please enter a question to query the knowledge base."
return rag_system.generate_response(query=query, k=k_chunks)
# Define the Gradio Interface
iface = gr.Interface(
fn=query_rag,
inputs=[
gr.Textbox(
label="Your Question",
placeholder="e.g., What is the purpose of the RAG architecture?",
lines=2
),
gr.Slider(
minimum=1,
maximum=5,
step=1,
value=3,
label="Number of Context Chunks (k)"
)
],
outputs=gr.Markdown(label="Grounded Answer (Powered by Gemini)"),
title="RAG Demo with Gemini 2.5 Flash and ChromaDB",
description=(
"Ask a question about Gemini 2.5 models or RAG architecture. "
"The system retrieves relevant context from its small knowledge base and uses **Gemini 2.5 Flash** "
"to generate a **grounded, cited answer**."
),
examples=[
["What makes the Gemini 2.5 Pro model suitable for long content?", 2],
["What is the role of ChromaDB in this setup and what embedding model is used?", 3],
["Who was the first person to walk on the moon?", 1]
]
)
# Launch the app
if __name__ == "__main__":
iface.launch()