import sys import warnings import pandas as pd import torch import gradio as gr import spaces from contextretriever import WikipediaContextRetriever from processor import MCQInferenceProcessor sys.unraisablehook = lambda unraisable: None warnings.filterwarnings("ignore", category=ResourceWarning) try: train_df = pd.read_csv("data/train.csv") test_df = pd.read_csv("data/test.csv") df_samples = pd.concat([train_df, test_df], ignore_index=True) except FileNotFoundError: print("⚠️ Warning: Sample CSV data files not found in 'data/' directory. Random button will be empty.") df_samples = pd.DataFrame() def get_random_question_from_local(): """Pulls a random question row from the CSV and returns values for the UI inputs.""" if df_samples.empty: return "No samples available", "", "", "", "", "" random_row = df_samples.sample(n=1).iloc[0] return ( str(random_row.get("prompt", "")), str(random_row.get("A", "")), str(random_row.get("B", "")), str(random_row.get("C", "")), str(random_row.get("D", "")), str(random_row.get("E", "")) ) processor = MCQInferenceProcessor(model_repo_id="balamaniansp/bigru-map3loss-lr5e-5-layers3-drop02-vocab10k-seq512-emb256-hidden256-fold3") retriever = WikipediaContextRetriever(sample_size=1500) @spaces.GPU def run_mcq_pipeline(question, opt_a, opt_b, opt_c, opt_d, opt_e): # Fallback to avoid empty string matrix processing errors options_dict = { "A": opt_a if opt_a.strip() else "[Empty Option A]", "B": opt_b if opt_b.strip() else "[Empty Option B]", "C": opt_c if opt_c.strip() else "[Empty Option C]", "D": opt_d if opt_d.strip() else "[Empty Option D]", "E": opt_e if opt_e.strip() else "[Empty Option E]" } # Context extraction through RAG framework retrieved_blocks = retriever._get_context(question, top_n=1) if retrieved_blocks and len(retrieved_blocks) > 0: context_str = retrieved_blocks[0].get("text", "No context found.") context_source = retrieved_blocks[0].get("title", "Unknown Source") else: context_str = "No applicable Wikipedia context retrieved." context_source = "Fallback System" # Model inference call execution probabilities = processor.predict(context=context_str, question=question, options=options_dict) context_feedback = f"📚 Source: {context_source}\n\nFetched Context:\n{context_str}" return probabilities, context_feedback # Layout definition using standard Gradio Blocks with gr.Blocks(theme=gr.themes.Soft()) as demo: gr.Markdown("# 🧠 Academic Presentation: Open-Domain MCQ Solver (RAG)") gr.Markdown("Type a question manually, or click **🎲 Load Random Sample Question** to test an existing dataset question.") with gr.Row(): with gr.Column(scale=2): input_q = gr.Textbox(lines=2, label="Question Stem", placeholder="Enter your question here...") input_a = gr.Textbox(label="Option A") input_b = gr.Textbox(label="Option B") input_c = gr.Textbox(label="Option C") input_d = gr.Textbox(label="Option D") input_e = gr.Textbox(label="Option E") with gr.Row(): sample_btn = gr.Button("🎲 Load Random Sample Question", variant="secondary") submit_btn = gr.Button("🚀 Run Inference", variant="primary") with gr.Column(scale=2): output_labels = gr.Label(label="Model Probability Confidence Distribution") output_context = gr.Textbox(label="Retrieved Reference Framework Data", interactive=False, lines=10) # Wire up action dependencies submit_btn.click( fn=run_mcq_pipeline, inputs=[input_q, input_a, input_b, input_c, input_d, input_e], outputs=[output_labels, output_context] ) sample_btn.click( fn=get_random_question_from_local, inputs=[], outputs=[input_q, input_a, input_b, input_c, input_d, input_e] ) if __name__ == "__main__": demo.launch(server_name="0.0.0.0", server_port=7860, prevent_thread_lock=False)