| import time |
| import gradio as gr |
|
|
| import torch |
| from transformers import GPT2LMHeadModel, GPT2Tokenizer |
| from peft import PeftModel, LoraConfig |
|
|
|
|
| tokenizer = GPT2Tokenizer.from_pretrained("gpt2-medium") |
| tokenizer.pad_token = tokenizer.eos_token |
| base_model = GPT2LMHeadModel.from_pretrained("gpt2-medium") |
|
|
| lora_config = LoraConfig( |
| r=8, |
| lora_alpha=16, |
| target_modules=["c_fc", "c_proj", "c_attn"], |
| lora_dropout=0.1, |
| task_type="CAUSAL_LM" |
| ) |
| |
| finetuned_model = PeftModel.from_pretrained(base_model, "./lora_ft_weights", config=lora_config) |
| finetuned_model.eval() |
|
|
|
|
|
|
| |
| |
| |
| def qa_system(method, question): |
| start_time = time.time() |
|
|
| if not question.strip(): |
| return "**Error:** Please enter a question.", 0.0, "0 seconds", "" |
|
|
| |
| if method == "Retrieval-Augmented Generation (RAG)": |
| answer = "Using RAG: Based on retrieved financial documents, the answer is $95,000,000." |
| model_name = "RAG-based Model" |
| confidence = 0.92 |
| else: |
| model_name = "GPT 2-finetuned" |
| |
| financial_keywords = [ |
| 'revenue', 'profit', 'earnings', 'financial', 'income', 'balance', |
| 'cash', 'debt', 'equity', 'assets', 'market', 'investment', 'sales', |
| 'cost', 'margin', 'growth', 'compliance', 'risk', 'customer' |
| ] |
| for text in question: |
| |
| text_lower = text.lower() |
| if any(pattern in text_lower for pattern in financial_keywords): |
| return "This question does not seem to be related to Finance" |
|
|
| prompt = f"You are a financial assistant.\nUse the context below to answer the question.\n\nQuestion: {question}\nAnswer:" |
| |
| inputs = tokenizer(prompt, return_tensors="pt", truncation=True, max_length=512) |
| |
| with torch.no_grad(): |
| outputs = finetuned_model.generate( |
| **inputs, |
| max_length=inputs['input_ids'].shape[1] + 100, |
| temperature=0.7, |
| do_sample=True, |
| pad_token_id=tokenizer.eos_token_id, |
| ) |
| |
| generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True) |
| answer = generated_text.split("Answer:")[-1].strip() |
| |
| end_time = time.time() |
| response_time = round(end_time - start_time, 2) |
|
|
| return ( |
| f"**Method:** {model_name}", |
| 0.95, |
| f"{response_time} seconds", |
| answer |
| ) |
|
|
| |
| |
| |
| with gr.Blocks(css=""" |
| .radio-vertical .wrap { |
| flex-direction: column !important; |
| } |
| .radio-vertical .wrap > label { |
| margin-bottom: 8px !important; |
| margin-right: 0 !important; |
| } |
| .small-btn { |
| max-width: fit-content !important; |
| width: auto !important; |
| } |
| .small-btn button { |
| width: auto !important; |
| min-width: unset !important; |
| padding: 8px 16px !important; |
| font-size: 16px !important; |
| white-space: nowrap !important; |
| max-width: fit-content !important; |
| } |
| """) as demo: |
| gr.Markdown( |
| """ |
| # 📊 Comparative Financial QA System |
| An implementation comparing **Retrieval-Augmented Generation (RAG)** and a **Fine-Tuned on LoRA and Replay-Based Learning** GPT 2 model for answering questions on financial reports. |
| """ |
| ) |
|
|
| |
| method = gr.Radio( |
| choices=["Retrieval-Augmented Generation (RAG)", "Fine-Tuned Model"], |
| label="Choose QA Method:", |
| value="Fine-Tuned Model", |
| interactive=True, |
| elem_classes="radio-vertical" |
| ) |
|
|
| |
| question = gr.Textbox( |
| label="Ask a question about Nice's 2023-2024 financials:", |
| placeholder="e.g., What was the total revenue in 2023?" |
| ) |
|
|
| |
| submit_btn = gr.Button("Get Answer", elem_classes="small-btn") |
|
|
| |
| with gr.Group(visible=False) as output_section: |
| method_output = gr.Markdown() |
| confidence_output = gr.Number(label="Model Confidence") |
| response_time_output = gr.Textbox(label="Response Time") |
| answer_output = gr.Markdown(label="Answer") |
|
|
| |
| def handle_submit(method_val, question_val): |
| |
| results = qa_system(method_val, question_val) |
| return [gr.Group(visible=True)] + list(results) |
| |
| submit_btn.click( |
| handle_submit, |
| inputs=[method, question], |
| outputs=[output_section, method_output, confidence_output, response_time_output, answer_output] |
| ) |
|
|
| |
| |
| |
| if __name__ == "__main__": |
| demo.launch() |
|
|