Bajaj_Hackathon_Space / workflow.py
Abhirup073's picture
changed to google embeddings
309744b
Raw
History Blame Contribute Delete
5.9 kB
from typing import TypedDict, List
from langchain_core.documents import Document
from langgraph.graph import StateGraph, END
from langchain_google_genai import ChatGoogleGenerativeAI
from langchain_core.prompts import ChatPromptTemplate
from langchain.retrievers import EnsembleRetriever
from models import Question, FinalAnswer, GeneratedQueries
from config import ANSWER_LLM_MODEL, QUERY_LLM_MODEL, GOOGLE_API_KEY
class GraphState(TypedDict):
original_questions: List[Question]
decomposed_questions: GeneratedQueries
retriever: EnsembleRetriever
documents: List[List[Document]]
answers: List[FinalAnswer]
class RAGWorkflow:
def __init__(self):
self.generation_llm = ChatGoogleGenerativeAI(model=ANSWER_LLM_MODEL, google_api_key=GOOGLE_API_KEY, temperature=0.0)
self.decomposition_llm = ChatGoogleGenerativeAI(model=QUERY_LLM_MODEL, google_api_key=GOOGLE_API_KEY, temperature=0.0)
self.graph = self._build_graph()
def _query_decomposition_node(self, state: GraphState):
prompt = ChatPromptTemplate.from_template(
"""### ROLE: You are a world-class expert in information retrieval.
### GOAL: 🎯 For each user question, generate exactly 3 distinct, self-contained search queries.
### πŸ“œ INSTRUCTIONS:
- **Multi-Angle Approach:** Generate queries covering: 1. Definitional, 2. Procedural/Conditional, and 3. Quantitative/Limit.
- **Use Policy Jargon:** Use precise insurance terminology.
- **Be Self-Contained:** Each query must be understandable on its own.
### OUTPUT FORMAT: You MUST return a Pydantic `GeneratedQueries` object.
### USER QUESTIONS:
{questions}"""
)
questions_str = "\n".join(f"{i+1}. {q.question}" for i, q in enumerate(state["original_questions"]))
structured_llm = self.decomposition_llm.with_structured_output(GeneratedQueries)
decomposition_chain = prompt | structured_llm
generated_lists: GeneratedQueries = decomposition_chain.invoke({"questions": questions_str})
for i, el in enumerate(generated_lists.lst):
el.queries.append(state["original_questions"][i].question)
return {"decomposed_questions": generated_lists}
def _retrieval_node(self, state: GraphState):
all_queries = [q for query_list in state["decomposed_questions"].lst for q in query_list.queries]
retrieved_docs_lists = state["retriever"].batch(all_queries)
queries_per_question = len(state["decomposed_questions"].lst[0].queries)
final_documents: List[List[Document]] = []
for i in range(len(state["original_questions"])):
start_index, end_index = i * queries_per_question, (i + 1) * queries_per_question
single_question_docs = [doc for docs_list in retrieved_docs_lists[start_index:end_index] for doc in docs_list]
unique_docs = {doc.page_content: doc for doc in single_question_docs}
final_documents.append(list(unique_docs.values()))
return {"documents": final_documents}
def _generation_node(self, state: GraphState):
prompt_template = """### ROLE: You are a meticulous, AI-powered Insurance Claims Adjudicator.
### πŸ“œ CRITICAL MANDATES:
1. **Context is Absolute:** Your answer MUST be derived *only* from the provided CONTEXT.
2. **Exhaustive Factual Extraction:** Extract **all** relevant details: conditions, waiting periods, monetary limits, sub-limits, percentages, and exclusions.
3. **Handle Missing Information:** If context is insufficient, state: "The provided context does not contain sufficient information to answer this question."
4. **Be Factual:** Stick to the facts. No pleasantries.
### OUTPUT FORMAT: You MUST respond with a Pydantic `FinalAnswer` object.
### CONTEXT:
{context}
### QUESTION:
{question}
"""
final_answers = []
questions = state["original_questions"]
documents = state["documents"]
prompt = ChatPromptTemplate.from_template(prompt_template)
structured_llm = self.generation_llm.with_structured_output(FinalAnswer)
generation_chain = prompt | structured_llm
for i, question in enumerate(questions):
context_str = "\n\n---\n\n".join([doc.page_content for doc in documents[i]])
try:
single_answer = generation_chain.invoke({"context": context_str, "question": question.question})
if single_answer is None:
final_answers.append(FinalAnswer(answer="The answer could not be generated, likely due to a safety policy violation for this specific question."))
else:
final_answers.append(single_answer)
except Exception as e:
print(f"Error processing question '{question.question}': {e}")
final_answers.append(FinalAnswer(answer=f"An error occurred while processing this question."))
return {"answers": final_answers}
def _build_graph(self):
workflow = StateGraph(GraphState)
workflow.add_node("decompose_query", self._query_decomposition_node)
workflow.add_node("retrieve", self._retrieval_node)
workflow.add_node("generate", self._generation_node)
workflow.set_entry_point("decompose_query")
workflow.add_edge("decompose_query", "retrieve")
workflow.add_edge("retrieve", "generate")
workflow.add_edge("generate", END)
return workflow.compile()
def invoke(self, questions: List[Question], retriever: EnsembleRetriever) -> List[FinalAnswer]:
initial_state = {"original_questions": questions, "retriever": retriever}
final_state = self.graph.invoke(initial_state)
return final_state.get("answers", [])