Spaces:
Running on Zero
Running on Zero
| """ | |
| 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: | |
| 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 βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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;"> | |
| βΈ 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 Β· DeBERTa-v3 Β· RoBERTa-base Β· 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() |