QSolver / app.py
a-01a's picture
Create app.py
31edbfc verified
Raw History Blame Contribute Delete
7.85 kB
import spaces
import os
import torch
import numpy as np
import gradio as gr
from scipy.stats import rankdata
from scipy.special import softmax
from transformers import (
AutoModelForCausalLM,
AutoModelForMultipleChoice,
AutoModelForSequenceClassification,
AutoTokenizer,
BitsAndBytesConfig
)
from peft import PeftModel
HF_TOKEN = os.environ.get("HF_TOKEN")
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
FOLDS = [1]
BASE_MODEL_NAME_QWEN = "Qwen/Qwen3-4B-Instruct-2507"
MODEL1_RAG_QWEN3_REPO = "a-01a/QSolver_Decoder_V16"
MODEL2_DEBERTA_REPO = "a-01a/QSolver_Encoder_V2"
MODEL3_SCRATCH_REPO = "a-01a/QSolver_Scratch"
OPTION_LETTERS = ["A", "B", "C", "D", "E"]
IDX_TO_LABEL = {0: 'A', 1: 'B', 2: 'C', 3: 'D', 4: 'E'}
models_loaded = False
tokenizer_m1 = None
model1 = None
tokenizer_m2 = None
model2 = None
tokenizer_m3 = None
model3 = None
qwen_option_token_ids = None
def sharpen_probs_single(probs, power=1.5):
p_pow = np.power(probs, power)
return p_pow / np.sum(p_pow)
def to_percentile_ranks_single(probs):
return rankdata(probs) / len(probs)
def format_qwen3_prompt(context, question, opts):
ctx_line = f"Context: {context.strip()}\n" if context.strip() else ""
return (
f"<|im_start|>system\n"
f"You are a scientific expert. Base your answer STRICTLY on the provided Context. "
f"Output ONLY the single letter corresponding to the correct option (A, B, C, D, or E).<|im_end|>\n"
f"<|im_start|>user\n"
f"{ctx_line}Question: {question}\n"
f"A) {opts['A']}\nB) {opts['B']}\nC) {opts['C']}\nD) {opts['D']}\nE) {opts['E']}<|im_end|>\n"
f"<|im_start|>assistant\n"
)
def get_answer_token_id(tokenizer, letter):
probe = "<|im_start|>assistant\n"
return tokenizer(probe + letter, add_special_tokens=False)["input_ids"][len(tokenizer(probe, add_special_tokens=False)["input_ids"]):][0]
def load_all_models():
global models_loaded, tokenizer_m1, model1, tokenizer_m2, model2, tokenizer_m3, model3, qwen_option_token_ids
if models_loaded:
return
device = "cuda"
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16
)
tokenizer_m1 = AutoTokenizer.from_pretrained(BASE_MODEL_NAME_QWEN, token=HF_TOKEN, use_fast=True)
tokenizer_m1.truncation_side = "left"
tokenizer_m1.padding_side = "left"
if tokenizer_m1.pad_token is None:
tokenizer_m1.pad_token = tokenizer_m1.eos_token
qwen_base = AutoModelForCausalLM.from_pretrained(
BASE_MODEL_NAME_QWEN,
quantization_config=bnb_config,
device_map="auto",
token=HF_TOKEN
)
model1 = PeftModel.from_pretrained(qwen_base, MODEL1_RAG_QWEN3_REPO, subfolder=f"fold_{FOLDS[0]}", token=HF_TOKEN)
model1.eval()
qwen_option_token_ids = [get_answer_token_id(tokenizer_m1, l) for l in OPTION_LETTERS]
tokenizer_m2 = AutoTokenizer.from_pretrained(MODEL2_DEBERTA_REPO, subfolder=f"fold_{FOLDS[0]}", token=HF_TOKEN, use_fast=True)
model2 = AutoModelForMultipleChoice.from_pretrained(
MODEL2_DEBERTA_REPO,
subfolder=f"fold_{FOLDS[0]}",
token=HF_TOKEN,
torch_dtype=torch.float16
).to(device)
model2.eval()
tokenizer_m3 = AutoTokenizer.from_pretrained(MODEL3_SCRATCH_REPO, subfolder=f"fold_{FOLDS[0]}", token=HF_TOKEN, use_fast=True)
model3 = AutoModelForSequenceClassification.from_pretrained(
MODEL3_SCRATCH_REPO,
subfolder=f"fold_{FOLDS[0]}",
token=HF_TOKEN,
torch_dtype=torch.float16,
num_labels=1
).to(device)
model3.eval()
models_loaded = True
@spaces.GPU
@torch.inference_mode()
def run_inference(context, question, opt_a, opt_b, opt_c, opt_d, opt_e):
load_all_models()
device = "cuda"
opts = {"A": opt_a, "B": opt_b, "C": opt_c, "D": opt_d, "E": opt_e}
qwen_prompt = format_qwen3_prompt(context, question, opts)
inputs_m1 = tokenizer_m1([qwen_prompt], return_tensors="pt", padding=True, truncation=True, max_length=1280).to(device)
out_m1 = model1(**inputs_m1)
logits_m1 = out_m1.logits[:, -1, qwen_option_token_ids].float().cpu().numpy()[0]
m1_probs = softmax(logits_m1)
q_repeats = [question] * 5
opt_list = [opts["A"], opts["B"], opts["C"], opts["D"], opts["E"]]
inputs_m2 = tokenizer_m2(q_repeats, opt_list, padding=True, truncation=True, max_length=512, return_tensors="pt")
inputs_m2 = {k: v.unsqueeze(0).to(device) for k, v in inputs_m2.items()}
out_m2 = model2(**inputs_m2)
m2_probs = softmax(out_m2.logits.cpu().numpy()[0])
scratch_texts = [f"Question: {question}\nOption: {opt}" for opt in opt_list]
inputs_m3 = tokenizer_m3(scratch_texts, padding=True, truncation=True, max_length=256, return_tensors="pt").to(device)
out_m3 = model3(**inputs_m3)
logits_m3 = out_m3.logits.squeeze(-1).cpu().numpy()
m3_probs = softmax(logits_m3)
m1_sharp = sharpen_probs_single(m1_probs, power=1.5)
m2_sharp = sharpen_probs_single(m2_probs, power=1.8)
m3_sharp = sharpen_probs_single(m3_probs, power=1.5)
m1_ranks = to_percentile_ranks_single(m1_probs)
m2_ranks = to_percentile_ranks_single(m2_probs)
m3_ranks = to_percentile_ranks_single(m3_probs)
deb_sorted = np.sort(m2_probs)[::-1]
margin = deb_sorted[0] - deb_sorted[1]
if margin >= 0.35:
w_deb, w_scratch, w_qwen = 0.82, 0.11, 0.07
elif margin >= 0.15:
w_deb, w_scratch, w_qwen = 0.62, 0.23, 0.15
else:
w_deb, w_scratch, w_qwen = 0.42, 0.34, 0.24
p_blend = (w_deb * m2_sharp) + (w_scratch * m3_sharp) + (w_qwen * m1_sharp)
r_blend = (w_deb * m2_ranks) + (w_scratch * m3_ranks) + (w_qwen * m1_ranks)
final_blend_scores = (0.85 * p_blend) + (0.15 * r_blend)
top3_idx = np.argsort(-final_blend_scores)[:3]
top3_labels = " ".join([IDX_TO_LABEL[i] for i in top3_idx])
def dict_format(probs):
return {IDX_TO_LABEL[i]: float(probs[i]) for i in range(5)}
return dict_format(m1_probs), dict_format(m2_probs), dict_format(m3_probs), top3_labels
with gr.Blocks(title="TriModel Optimized Ensemble", theme=gr.themes.Soft()) as demo:
gr.Markdown("TriModel Optimized MCQ Solver Ensemble")
gr.Markdown("Input a question and 5 options. Optionally, provide context for Qwen3's RAG component. The application dynamically weights the 3 models based on DeBERTa's prediction confidence margin.")
with gr.Row():
with gr.Column(scale=2):
context_in = gr.Textbox(label="Context (Optional RAG input)", lines=3, placeholder="Paste background text here...")
question_in = gr.Textbox(label="Question", lines=2, placeholder="What is the process of...")
opt_a_in = gr.Textbox(label="Option A")
opt_b_in = gr.Textbox(label="Option B")
opt_c_in = gr.Textbox(label="Option C")
opt_d_in = gr.Textbox(label="Option D")
opt_e_in = gr.Textbox(label="Option E")
submit_btn = gr.Button("Run Inference", variant="primary")
with gr.Column(scale=3):
gr.Markdown("Ensemble Final Prediction (Map@3)")
ensemble_out = gr.Textbox(label="Top 3 Predictions", text_align="center", scale=2)
gr.Markdown("Individual Model Probabilities")
with gr.Row():
m1_out = gr.Label(label="Model 1: Qwen3 RAG (4B)")
m2_out = gr.Label(label="Model 2: DeBERTa-v3 Encoder")
m3_out = gr.Label(label="Model 3: Scratch Classifier")
submit_btn.click(
fn=run_inference,
inputs=[context_in, question_in, opt_a_in, opt_b_in, opt_c_in, opt_d_in, opt_e_in],
outputs=[m1_out, m2_out, m3_out, ensemble_out]
)
if __name__ == "__main__":
demo.launch()