madhur71/Multi_Agent_Workflow
0
1 2import gradio as gr3import numpy as np4import faiss5import torch6 7from sentence_transformers import SentenceTransformer8from transformers import (9 AutoTokenizer,10 AutoModelForSeq2SeqLM,11 AutoModelForQuestionAnswering,12 AutoModelForSequenceClassification,13 AutoModelForCausalLM,14)15 16with open("documents.txt", "r", encoding="utf-8") as f:17 docs = f.read().split("\n")18 19docs = [doc.strip() for doc in docs if doc.strip()]20 21embed = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")22 23doc_emb = embed.encode(docs)24doc_emb = np.array(doc_emb).astype("float32")25 26index = faiss.IndexFlatL2(doc_emb.shape[1])27index.add(doc_emb)28 29def retrieve_docs(query, k=2):30 query_emb = embed.encode([query])31 32 distances, indices = index.search(33 np.array(query_emb).astype("float32"),34 k35 )36 37 retrieved_docs = [docs[idx] for idx in indices[0]]38 39 return retrieved_docs40 41device = torch.device(42 "cuda" if torch.cuda.is_available() else "cpu"43)44 45sum_model_name = "facebook/bart-large-cnn"46 47sum_tokenizer = AutoTokenizer.from_pretrained(48 sum_model_name49)50 51sum_model = AutoModelForSeq2SeqLM.from_pretrained(52 sum_model_name53).to(device)54 55qa_model_name = "distilbert-base-cased-distilled-squad"56 57qa_tokenizer = AutoTokenizer.from_pretrained(58 qa_model_name59)60 61qa_model = AutoModelForQuestionAnswering.from_pretrained(62 qa_model_name63).to(device)64 65senti_model_name = (66 "distilbert-base-uncased-finetuned-sst-2-english"67)68 69senti_tokenizer = AutoTokenizer.from_pretrained(70 senti_model_name71)72 73senti_model = AutoModelForSequenceClassification.from_pretrained(74 senti_model_name75).to(device)76 77gen_model_name = "gpt2"78 79gen_tokenizer = AutoTokenizer.from_pretrained(80 gen_model_name81)82 83gen_model = AutoModelForCausalLM.from_pretrained(84 gen_model_name85).to(device)86 87if gen_tokenizer.pad_token is None:88 gen_tokenizer.pad_token = gen_tokenizer.eos_token89 90def summarize(text):91 92 inputs = sum_tokenizer(93 text,94 return_tensors="pt",95 truncation=True,96 max_length=1024,97 ).to(device)98 99 summary_ids = sum_model.generate(100 inputs["input_ids"],101 attention_mask=inputs["attention_mask"],102 max_length=60,103 min_length=20,104 do_sample=False,105 )106 107 summary = sum_tokenizer.decode(108 summary_ids[0],109 skip_special_tokens=True110 )111 112 return summary113 114def answer_question(question, context):115 116 inputs = qa_tokenizer(117 question,118 context,119 return_tensors="pt",120 truncation=True,121 max_length=512,122 ).to(device)123 124 outputs = qa_model(**inputs)125 126 start_idx = torch.argmax(127 outputs.start_logits128 ).item()129 130 end_idx = torch.argmax(131 outputs.end_logits132 ).item() + 1133 134 tokens = inputs["input_ids"][0][start_idx:end_idx]135 136 answer = qa_tokenizer.decode(137 tokens,138 skip_special_tokens=True139 )140 141 return answer142 143def sentiment_analysis(text):144 145 inputs = senti_tokenizer(146 text,147 return_tensors="pt",148 truncation=True,149 max_length=512,150 ).to(device)151 152 outputs = senti_model(**inputs)153 154 probs = torch.nn.functional.softmax(155 outputs.logits,156 dim=-1157 )158 159 pred = torch.argmax(probs).item()160 161 label = "POSITIVE" if pred == 1 else "NEGATIVE"162 163 score = probs[0][pred].item()164 165 return f"{label} ({score:.4f})"166 167def generate_text(prompt):168 169 inputs = gen_tokenizer(170 prompt,171 return_tensors="pt"172 ).to(device)173 174 outputs = gen_model.generate(175 **inputs,176 max_new_tokens=50,177 do_sample=True,178 temperature=0.7,179 pad_token_id=gen_tokenizer.eos_token_id,180 )181 182 generated = gen_tokenizer.decode(183 outputs[0],184 skip_special_tokens=True185 )186 187 return generated188 189def coordinator(query):190 191 retrieved_docs = retrieve_docs(query, k=2)192 193 context = " ".join(retrieved_docs)194 195 summary = summarize(context)196 197 qa_res = answer_question(query, context)198 199 sentiment = sentiment_analysis(context)200 201 generated = generate_text(202 f"Context: {context}\n"203 f"Question: {query}\n"204 f"Answer:"205 )206 207 retrieved_text = "\n\n".join(retrieved_docs)208 209 return (210 retrieved_text,211 summary,212 qa_res,213 sentiment,214 generated,215 )216 217custom_css = """218body {219 background: linear-gradient(220 135deg,221 #0f172a,222 #1e293b,223 #111827224 );225 font-family: Arial, sans-serif;226}227 228.gradio-container {229 background: rgba(17, 24, 39, 0.96) !important;230 border-radius: 20px;231 padding: 25px;232 box-shadow: 0 10px 35px rgba(0,0,0,0.5);233}234 235textarea, input {236 background-color: #1e293b !important;237 color: white !important;238 border-radius: 12px !important;239 border: 1px solid #334155 !important;240}241 242button {243 background: linear-gradient(244 90deg,245 #2563eb,246 #7c3aed247 ) !important;248 249 color: white !important;250 border: none !important;251 border-radius: 12px !important;252 font-weight: bold !important;253 transition: 0.3s ease;254}255 256button:hover {257 transform: scale(1.03);258 opacity: 0.95;259}260 261h1, h3 {262 text-align: center;263 color: white;264}265 266footer {267 visibility: hidden;268}269"""270 271with gr.Blocks(css=custom_css) as demo:272 273 gr.Markdown("""274 # Multi-Agent AI Assistant275 276 ### Retrieval + Summarization + QA + Intent + Generation277 """)278 279 query_input = gr.Textbox(280 lines=3,281 placeholder="Ask something...",282 label="Enter Query"283 )284 285 submit_btn = gr.Button(286 "Run Multi-Agent Pipeline"287 )288 289 with gr.Tab("Retrieved Documents"):290 retrieved_output = gr.Textbox(291 lines=10,292 label="Retrieved Context"293 )294 295 with gr.Tab("Summary Agent"):296 summary_output = gr.Textbox(297 lines=6,298 label="Summarized Output"299 )300 301 with gr.Tab("Question Answering Agent"):302 qa_output = gr.Textbox(303 lines=4,304 label="Answer"305 )306 307 with gr.Tab("Intent Detection Agent"):308 sentiment_output = gr.Textbox(309 lines=3,310 label="Detected Intent"311 )312 313 with gr.Tab("Generation Agent"):314 generated_output = gr.Textbox(315 lines=8,316 label="Generated Response"317 )318 319 submit_btn.click(320 fn=coordinator,321 inputs=query_input,322 outputs=[323 retrieved_output,324 summary_output,325 qa_output,326 sentiment_output,327 generated_output,328 ]329 )330 331demo.launch()332 