Smart-MCQ-App / app.py
udaypratap's picture
Create app.py
8a09357 verified
Raw
History Blame Contribute Delete
2.13 kB
import gradio as gr
import torch
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModelForMultipleChoice
# 1. Load fine-tuned model and tokenizer from Hugging Face Hub
MODEL_ID = "udaypratap/smart-mcq-solver"
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForMultipleChoice.from_pretrained(MODEL_ID)
model.eval()
# 2. Define prediction function
def predict_mcq(prompt, option_a, option_b, option_c, option_d, option_e):
options = [option_a, option_b, option_c, option_d, option_e]
labels = ["A", "B", "C", "D", "E"]
# Format inputs for AutoModelForMultipleChoice
first_sentences = [prompt] * 5
second_sentences = options
inputs = tokenizer(
first_sentences,
second_sentences,
truncation=True,
padding=True,
max_length=256,
return_tensors="pt"
)
# Reshape input tensors for multiple choice model: (batch_size=1, num_choices=5, seq_len)
input_ids = inputs["input_ids"].unsqueeze(0)
attention_mask = inputs["attention_mask"].unsqueeze(0)
with torch.no_grad():
outputs = model(input_ids=input_ids, attention_mask=attention_mask)
logits = outputs.logits.squeeze(0)
probs = F.softmax(logits, dim=-1)
# Get top 3 predicted choices
top3_indices = torch.topk(probs, k=3).indices.tolist()
top3_choices = [f"{labels[idx]} ({probs[idx].item():.2%})" for idx in top3_indices]
return f"Top 3 Predicted Answers: {', '.join(top3_choices)}"
# 3. Create Gradio Interface
demo = gr.Interface(
fn=predict_mcq,
inputs=[
gr.Textbox(label="Question Prompt", placeholder="Enter your question here..."),
gr.Textbox(label="Option A"),
gr.Textbox(label="Option B"),
gr.Textbox(label="Option C"),
gr.Textbox(label="Option D"),
gr.Textbox(label="Option E")
],
outputs=gr.Textbox(label="Predictions"),
title="Smart MCQ Solver",
description="Enter a question prompt along with 5 options to get the top predicted answers."
)
if __name__ == "__main__":
demo.launch()