CoolFace
Apppublic

madhur71/Multi_Agent_Workflow

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
app.py332 linesDownload Raw Back to root
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