jems9376's picture
Update app.py
d4dafc8 verified
Raw
History Blame Contribute Delete
18.2 kB
"""
Smart MCQ Solver β€” HuggingFace Spaces
Ensemble: CNN + DeBERTa-v3-small + RoBERTa-base
Loading strategy
----------------
- shared_vocab.json : loaded from local models/ folder in the Space repo
- CNN / DeBERTa / RoBERTa : all downloaded from jems9376/<repo> on HF Hub
For DeBERTa & RoBERTa:
1. snapshot_download pulls tokenizer files + fine-tuned .pt
2. config.json inside that snapshot tells us the BASE model name
3. TransformerMCQScorer loads BASE weights from HF Hub, then overlays the .pt
"""
import os, re, json
import numpy as np
import torch
import torch.nn as nn
import gradio as gr
from transformers import AutoTokenizer, AutoModel
from huggingface_hub import snapshot_download
# ── ZeroGPU support ───────────────────────────────────────────────────────────
try:
import spaces
ZERO_GPU = True
except ImportError:
ZERO_GPU = False
class spaces:
@staticmethod
def GPU(fn=None, duration=60):
return fn if fn else (lambda f: f)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Device: {DEVICE}")
MODELS_DIR = "models"
HF_CNN_REPO = "jems9376/cnn"
HF_DEBERTA_REPO = "jems9376/deberta_rag"
HF_ROBERTA_REPO = "jems9376/roberta_rag"
# ── CNN tokenisation helpers ──────────────────────────────────────────────────
TOKEN_RE = re.compile(r"[A-Za-z]+|\d+|[^\sA-Za-z\d]")
PAD, UNK, SEP = "<pad>", "<unk>", "<sep>"
def simple_tokenize(text):
return TOKEN_RE.findall(str(text).lower())
def build_first_sentence(context, prompt):
if context and str(context).strip():
return f"Context: {context} Question: {prompt}"
return f"Question: {prompt}"
def make_scratch_encode_fn(word2id, max_len):
def encode_pair(text_a, text_b):
toks = (simple_tokenize(text_a) + [SEP] + simple_tokenize(text_b))[:max_len]
ids = [word2id.get(t, word2id[UNK]) for t in toks]
mask = [1] * len(ids)
pad_len = max_len - len(ids)
ids += [word2id[PAD]] * pad_len
mask += [0] * pad_len
return ids, mask
return encode_pair
def make_hf_encode_fn(tokenizer, max_len):
def encode_pair(text_a, text_b):
enc = tokenizer(text_a, text_b, max_length=max_len,
padding="max_length", truncation=True)
return enc["input_ids"], enc["attention_mask"]
return encode_pair
# ── Model architectures ───────────────────────────────────────────────────────
class CNNScorer(nn.Module):
def __init__(self, vocab_size, emb_dim=128, num_filters=64,
kernel_sizes=(2, 3, 4), dropout=0.3):
super().__init__()
self.embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=0)
self.convs = nn.ModuleList([
nn.Conv1d(emb_dim, num_filters, k, padding=k // 2)
for k in kernel_sizes
])
self.dropout = nn.Dropout(dropout)
self.scorer = nn.Sequential(
nn.Linear(num_filters * len(kernel_sizes), 64), nn.ReLU(),
nn.Dropout(dropout), nn.Linear(64, 1),
)
def encode(self, input_ids, attn_mask):
x = self.embedding(input_ids).transpose(1, 2)
pooled = []
for conv in self.convs:
c = conv(x)
L = min(c.size(-1), attn_mask.size(-1))
c = c[..., :L]
m = attn_mask[:, :L].unsqueeze(1)
c = c.masked_fill(~m, float("-inf"))
pooled.append(c.max(dim=-1).values)
return self.dropout(torch.cat(pooled, dim=-1))
def forward(self, input_ids, attn_mask):
B, K, L = input_ids.shape
enc = self.encode(input_ids.view(B * K, L), attn_mask.view(B * K, L))
return self.scorer(enc).view(B, K)
class TransformerMCQScorer(nn.Module):
"""
Wraps a HuggingFace encoder + a linear scorer head.
base_model_name : HF model-hub ID for the BASE model (e.g. "microsoft/deberta-v3-small")
"""
def __init__(self, base_model_name: str, dropout: float = 0.1):
super().__init__()
self.encoder = AutoModel.from_pretrained(base_model_name)
h = self.encoder.config.hidden_size
self.drop = nn.Dropout(dropout)
self.scorer = nn.Linear(h, 1)
def forward(self, input_ids, attention_mask):
B, C, L = input_ids.shape
flat_ids = input_ids.view(B * C, L)
flat_mask = attention_mask.view(B * C, L)
out = self.encoder(input_ids=flat_ids, attention_mask=flat_mask,
return_dict=True)
cls = self.drop(out.last_hidden_state[:, 0, :])
return self.scorer(cls).view(B, C)
# ── Loaders ───────────────────────────────────────────────────────────────────
def _load_cnn(cnn_dir: str, vocab_path: str):
cfg = json.load(open(os.path.join(cnn_dir, "config.json")))
word2id = json.load(open(vocab_path))
m = CNNScorer(cfg["vocab_size"], cfg["emb_dim"],
cfg["num_filters"], cfg["kernel_sizes"]).to(DEVICE)
m.load_state_dict(torch.load(
os.path.join(cnn_dir, "cnn_best.pt"),
map_location=DEVICE, weights_only=True))
m.eval()
m.float()
return m, make_scratch_encode_fn(word2id, cfg["max_len"])
def _load_transformer(snapshot_dir: str, ckpt_filename: str):
"""
snapshot_dir : local path returned by snapshot_download (has tokenizer + .pt, NO base weights)
ckpt_filename : e.g. "deberta_rag_best.pt"
Steps:
1. Read config.json β†’ get base model name (e.g. "microsoft/deberta-v3-small")
2. Build TransformerMCQScorer from that base model (downloads base weights from HF Hub)
3. Load the fine-tuned .pt state dict on top
4. Load tokenizer from snapshot_dir
"""
cfg_path = os.path.join(snapshot_dir, "config.json")
cfg = json.load(open(cfg_path))
base_model_name = cfg["model_name"] # e.g. "microsoft/deberta-v3-small"
max_len = cfg["max_len"]
dropout = cfg.get("dropout", 0.1)
print(f" Base model: {base_model_name}")
# Build model with base weights from HF Hub
m = TransformerMCQScorer(base_model_name, dropout).to(DEVICE).float()
# Load fine-tuned weights on top
ckpt_path = os.path.join(snapshot_dir, ckpt_filename)
state = torch.load(ckpt_path, map_location=DEVICE, weights_only=True)
m.load_state_dict(state)
m.eval()
m.float() # ensure float32 even after ZeroGPU re-allocation
# Tokenizer lives in the snapshot (tokenizer.json + tokenizer_config.json)
tok = AutoTokenizer.from_pretrained(snapshot_dir)
return m, make_hf_encode_fn(tok, max_len)
def load_models():
loaded = {}
# ── CNN ───────────────────────────────────────────────────────────────────
# shared_vocab.json lives in the Space repo under models/
# The CNN checkpoint itself is always pulled from HF Hub
vocab_path = os.path.join(MODELS_DIR, "shared_vocab.json")
try:
print(f"Downloading CNN snapshot from {HF_CNN_REPO} …")
cnn_dir = snapshot_download(repo_id=HF_CNN_REPO)
m, enc_fn = _load_cnn(cnn_dir, vocab_path)
loaded["cnn"] = (m, enc_fn)
print("βœ“ CNN loaded")
except Exception as e:
print(f"βœ— CNN failed: {e}")
# ── DeBERTa ───────────────────────────────────────────────────────────────
# snapshot_download fetches ONLY the tokenizer files + .pt from jems9376/deberta_rag
# Base model weights are pulled separately via AutoModel.from_pretrained(base_model_name)
try:
print(f"Downloading DeBERTa snapshot from {HF_DEBERTA_REPO} …")
deb_snapshot = snapshot_download(repo_id=HF_DEBERTA_REPO)
m, enc_fn = _load_transformer(deb_snapshot, "deberta_rag_best.pt")
loaded["deberta_rag"] = (m, enc_fn)
print("βœ“ DeBERTa loaded")
except Exception as e:
print(f"βœ— DeBERTa failed: {e}")
# ── RoBERTa ───────────────────────────────────────────────────────────────
try:
print(f"Downloading RoBERTa snapshot from {HF_ROBERTA_REPO} …")
rob_snapshot = snapshot_download(repo_id=HF_ROBERTA_REPO)
m, enc_fn = _load_transformer(rob_snapshot, "roberta_rag_best.pt")
loaded["roberta_rag"] = (m, enc_fn)
print("βœ“ RoBERTa loaded")
except Exception as e:
print(f"βœ— RoBERTa failed: {e}")
if not loaded:
raise RuntimeError("No model checkpoints could be loaded.")
print(f"Active models: {list(loaded.keys())}")
return loaded
print("Loading models…")
MODELS = load_models()
print("All models ready.")
# ── Inference ─────────────────────────────────────────────────────────────────
@spaces.GPU(duration=120)
@torch.no_grad()
def predict(prompt, opt_a, opt_b, opt_c, opt_d, opt_e, context=""):
options = [opt_a, opt_b, opt_c, opt_d, opt_e]
opt_labels = ["A", "B", "C", "D", "E"]
first = build_first_sentence(context, prompt)
# ZeroGPU may re-allocate weights in fp16 between calls β€” force fp32 every time
for _, (model, _) in MODELS.items():
model.float()
all_logits = {}
for name, (model, enc_fn) in MODELS.items():
ids_all, mask_all = [], []
for opt in options:
ids, mask = enc_fn(first, str(opt))
ids_all.append(ids)
mask_all.append(mask)
input_ids = torch.tensor([ids_all], dtype=torch.long).to(DEVICE)
attn_mask = torch.tensor([mask_all], dtype=torch.bool).to(DEVICE)
logits = model(input_ids, attn_mask)
all_logits[name] = logits[0].float().cpu().numpy()
ensemble_logits = np.mean(list(all_logits.values()), axis=0)
ranked_idx = np.argsort(-ensemble_logits)
exp_l = np.exp(ensemble_logits - ensemble_logits.max())
probs = (exp_l / exp_l.sum() * 100) # percentages
# ── build HTML output ─────────────────────────────────────────────────────
rank_colors = ["#4f46e5", "#7c3aed", "#a855f7", "#d1d5db", "#d1d5db"]
rank_labels = ["1st", "2nd", "3rd", "4th", "5th"]
rows_html = ""
for rank, idx in enumerate(ranked_idx):
label = opt_labels[idx]
text = options[idx] if options[idx].strip() else "β€”"
pct = probs[idx]
bar_w = max(pct, 2) # min 2% so bar is always visible
color = rank_colors[rank]
is_top = rank == 0
bg = "background:#f5f3ff;border:1.5px solid #4f46e5;" if is_top else "background:#fafafa;border:1px solid #e5e7eb;"
weight = "font-weight:700;" if is_top else "font-weight:500;"
crown = "<span style='margin-right:6px;font-size:1rem;'>πŸ‘‘</span>" if is_top else ""
rows_html += f"""
<div style="margin-bottom:10px;padding:12px 14px;border-radius:10px;{bg}">
<div style="display:flex;align-items:center;justify-content:space-between;margin-bottom:6px;">
<div style="display:flex;align-items:center;gap:8px;">
<span style="display:inline-block;min-width:28px;padding:2px 7px;border-radius:5px;
background:{color};color:#fff;font-size:0.78rem;font-weight:700;text-align:center;">
{label}
</span>
{crown}<span style="{weight}font-size:0.95rem;color:#111;">{text}</span>
</div>
<span style="font-size:0.85rem;color:{color};font-weight:700;white-space:nowrap;margin-left:12px;">
{pct:.1f}%
</span>
</div>
<div style="height:6px;border-radius:99px;background:#e5e7eb;overflow:hidden;">
<div style="height:100%;width:{bar_w:.1f}%;border-radius:99px;background:{color};
transition:width 0.4s ease;"></div>
</div>
</div>"""
# per-model breakdown β€” convert logits β†’ softmax %, sort by rank per model
model_display = {"cnn": "CNN", "deberta_rag": "DeBERTa", "roberta_rag": "RoBERTa"}
model_rows_html = ""
for name, logits in all_logits.items():
exp_m = np.exp(logits - logits.max())
m_probs = exp_m / exp_m.sum() * 100 # softmax % for this model
m_rank = np.argsort(-m_probs) # sorted indices best→worst
top1_i = m_rank[0]
cells = ""
for rank_pos, idx in enumerate(m_rank):
is_top = rank_pos == 0
bg = "background:#ede9fe;" if is_top else ""
color = "color:#4f46e5;font-weight:700;" if is_top else "color:#6b7280;"
cells += (
f"<span style='display:inline-flex;align-items:center;gap:5px;"
f"padding:3px 8px;border-radius:6px;{bg}margin:2px;white-space:nowrap;'>"
f"<span style='font-weight:700;{color}'>{opt_labels[idx]}</span>"
f"<span style='{color}font-size:0.85em;'>{m_probs[idx]:.1f}%</span>"
f"</span>"
)
model_rows_html += f"""
<div style="display:flex;align-items:center;padding:8px 0;border-bottom:1px solid #f3f4f6;">
<span style="min-width:72px;font-weight:600;font-size:0.85rem;color:#374151;">
{model_display.get(name, name)}
</span>
<div style="display:flex;flex-wrap:wrap;gap:2px;">{cells}</div>
</div>"""
html = f"""
<div style="font-family:'Inter',sans-serif;max-width:700px;">
<p style="font-size:0.75rem;color:#9ca3af;margin:0 0 10px 0;letter-spacing:0.04em;text-transform:uppercase;">
Ensemble ranking
</p>
{rows_html}
<details style="margin-top:16px;">
<summary style="cursor:pointer;font-size:0.82rem;color:#6b7280;user-select:none;
list-style:none;display:flex;align-items:center;gap:6px;padding:4px 0;">
β–Έ &nbsp;What each model thinks
</summary>
<div style="margin-top:8px;padding:10px 12px;border-radius:10px;background:#fafafa;
border:1px solid #e5e7eb;">
<p style="font-size:0.75rem;color:#9ca3af;margin:0 0 8px 0;">
Ranked best β†’ worst per model, shown as confidence %
</p>
{model_rows_html}
</div>
</details>
</div>
"""
return html
# ── UI ────────────────────────────────────────────────────────────────────────
CSS = """
.gradio-container { max-width: 820px !important; margin: auto; }
footer { display: none !important; }
#predict-btn { min-width: 110px; }
"""
_theme = gr.themes.Default(
primary_hue="indigo",
font=[gr.themes.GoogleFont("Inter"), "sans-serif"],
)
with gr.Blocks(title="MCQ Solver", theme=_theme, css=CSS) as demo:
gr.Markdown(
"## MCQ Solver\n"
"<span style='color:#6b7280;font-size:0.95rem;'>"
"Ensemble Β· CNN &nbsp;Β·&nbsp; DeBERTa-v3 &nbsp;Β·&nbsp; RoBERTa-base &nbsp;Β·&nbsp; RAG-augmented"
"</span>"
)
with gr.Row(equal_height=False):
with gr.Column(scale=5):
prompt_box = gr.Textbox(
label="Question", lines=2,
placeholder="Type or paste your question here…"
)
with gr.Row():
opt_a = gr.Textbox(label="A", scale=1)
opt_b = gr.Textbox(label="B", scale=1)
with gr.Row():
opt_c = gr.Textbox(label="C", scale=1)
opt_d = gr.Textbox(label="D", scale=1)
opt_e = gr.Textbox(label="E", placeholder="Optional β€” leave blank if not used")
context_box = gr.Textbox(
label="Context (optional β€” paste a relevant passage for better accuracy)",
placeholder="e.g. a Wikipedia paragraph related to the question…",
lines=3,
)
predict_btn = gr.Button("Predict", variant="primary", elem_id="predict-btn")
output = gr.HTML()
predict_btn.click(
fn=predict,
inputs=[prompt_box, opt_a, opt_b, opt_c, opt_d, opt_e, context_box],
outputs=output,
)
gr.Examples(
examples=[
[
"What is the powerhouse of the cell?",
"Mitochondria", "Nucleus", "Ribosome", "Golgi apparatus", "Lysosome",
"Mitochondria generate most of the cell's supply of ATP, used as a source of chemical energy.",
],
[
"Which element has atomic number 1?",
"Hydrogen", "Helium", "Lithium", "Carbon", "Oxygen",
"Hydrogen is a chemical element with symbol H and atomic number 1.",
],
[
"Who wrote the play Hamlet?",
"Charles Dickens", "William Shakespeare", "Leo Tolstoy", "Mark Twain", "Homer",
"Hamlet is a tragedy written by William Shakespeare around 1600–1601.",
],
],
inputs=[prompt_box, opt_a, opt_b, opt_c, opt_d, opt_e, context_box],
outputs=output,
fn=predict,
cache_examples=False,
)
if __name__ == "__main__":
demo.launch()