Spaces:
Running
Running
| """ | |
| model.py β Model loading and inference logic for the Smart MCQ Solver. | |
| Pure Python module with no Streamlit dependency, so it can be used in | |
| notebooks, scripts, or any other context. | |
| Uses only the RoBERTa-base model (DeBERTa was found to be non-functional | |
| after evaluation on training data β 17% accuracy vs RoBERTa's 99%). | |
| """ | |
| import os | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| import re | |
| import json | |
| from transformers import AutoTokenizer, AutoModelForSequenceClassification | |
| from peft import PeftModel | |
| from .config import ( | |
| DEVICE, NUM_LABELS, LABEL2ID, ID2LABEL, MAX_LEN, | |
| MODEL_SOURCE, HF_TOKEN, | |
| ROBERTA_BASE_CHECKPOINT, | |
| ROBERTA_LOCAL_PATH, | |
| ROBERTA_HUB_REPO, | |
| ) | |
| # ββ Adapter config sanitisation ββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _sanitize_adapter_config(adapter_path: str): | |
| """Remove newer PEFT config fields that break older peft library versions.""" | |
| config_file = os.path.join(adapter_path, "adapter_config.json") | |
| if os.path.isfile(config_file): | |
| try: | |
| with open(config_file, "r") as f: | |
| cfg = json.load(f) | |
| valid_keys = { | |
| "r", "target_modules", "lora_alpha", "lora_dropout", "fan_in_fan_out", | |
| "bias", "use_rslora", "modules_to_save", "init_lora_weights", | |
| "layers_to_transform", "layers_pattern", "rank_pattern", "alpha_pattern", | |
| "megatron_config", "megatron_core", "loftq_config", "use_dora", | |
| "layer_replication", | |
| "peft_type", "auto_mapping", "base_model_name_or_path", "revision", | |
| "task_type", "inference_mode", | |
| } | |
| keys_to_remove = [k for k in cfg if k not in valid_keys] | |
| if keys_to_remove: | |
| for k in keys_to_remove: | |
| cfg.pop(k) | |
| with open(config_file, "w") as f: | |
| json.dump(cfg, f, indent=2) | |
| except Exception: | |
| pass | |
| # ββ Model loading ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def load_model(base_checkpoint: str, adapter_source: str): | |
| """ | |
| Load a base HuggingFace model and attach a PEFT LoRA adapter. | |
| Args: | |
| base_checkpoint: HuggingFace model ID (e.g. "roberta-base"). | |
| adapter_source: Either a local directory path or a HuggingFace Hub repo ID. | |
| Returns: | |
| (model, tokenizer) tuple, ready for inference on DEVICE. | |
| """ | |
| hub_kwargs = {"token": HF_TOKEN} if HF_TOKEN else {} | |
| if os.path.isdir(adapter_source): | |
| _sanitize_adapter_config(adapter_source) | |
| # Load tokenizer from the base checkpoint | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| base_checkpoint, use_fast=True, **hub_kwargs | |
| ) | |
| # Load the frozen base model | |
| base_model = AutoModelForSequenceClassification.from_pretrained( | |
| base_checkpoint, | |
| num_labels=5, | |
| **hub_kwargs, | |
| ) | |
| # Attach LoRA adapters | |
| try: | |
| model = PeftModel.from_pretrained(base_model, adapter_source, **hub_kwargs) | |
| except TypeError: | |
| # Fallback in case of PEFT version incompatibility | |
| from peft import LoraConfig | |
| config = LoraConfig.from_pretrained(adapter_source, **hub_kwargs) | |
| model = PeftModel(base_model, config) | |
| model.load_adapter(adapter_source, "default") | |
| model.to(DEVICE).eval() | |
| return model, tokenizer | |
| def load_roberta_model(): | |
| """ | |
| Load the RoBERTa model based on MODEL_SOURCE config. | |
| Returns: | |
| (roberta_model, roberta_tok) | |
| """ | |
| if MODEL_SOURCE == "hub": | |
| rob_source = ROBERTA_HUB_REPO | |
| print(f"Loading RoBERTa from HuggingFace Hub...") | |
| print(f" RoBERTa: {rob_source}") | |
| else: | |
| rob_source = ROBERTA_LOCAL_PATH | |
| print(f"Loading RoBERTa from local path...") | |
| print(f" RoBERTa: {rob_source}") | |
| if not os.path.isdir(rob_source): | |
| raise FileNotFoundError( | |
| f"RoBERTa adapter not found at: {rob_source}\n" | |
| f"Run the finetuning notebook first, or set MODEL_SOURCE=hub." | |
| ) | |
| rob_model, rob_tok = load_model(ROBERTA_BASE_CHECKPOINT, rob_source) | |
| print(f" β RoBERTa loaded on {DEVICE}") | |
| return rob_model, rob_tok | |
| # ββ Inference ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def get_probs(model, tokenizer, prompt: str, options: dict) -> np.ndarray: | |
| """ | |
| Run a forward pass and return a length-5 softmax probability array. | |
| Args: | |
| model: A PEFT-wrapped sequence classification model. | |
| tokenizer: The corresponding tokenizer. | |
| prompt: Question text. | |
| options: Dict {"A": ..., "B": ..., "C": ..., "D": ..., "E": ...}. | |
| Returns: | |
| numpy array of shape (5,) with probabilities [P(A)...P(E)]. | |
| """ | |
| text = ( | |
| f"{prompt}\n" | |
| f"A) {options['A']}\nB) {options['B']}\nC) {options['C']}\n" | |
| f"D) {options['D']}\nE) {options['E']}" | |
| ) | |
| encoded = tokenizer( | |
| text, | |
| truncation=True, max_length=MAX_LEN, | |
| padding="max_length", return_tensors="pt" | |
| ).to(DEVICE) | |
| logits = model(**encoded).logits.squeeze(0).float() | |
| return F.softmax(logits, dim=-1).cpu().numpy() | |
| def predict_mcq( | |
| prompt: str, | |
| options: dict, | |
| rob_model, rob_tok, | |
| ) -> dict: | |
| """ | |
| End-to-end MCQ prediction using RoBERTa. | |
| Args: | |
| prompt: Question text. | |
| options: Dict {"A": ..., "B": ..., "C": ..., "D": ..., "E": ...}. | |
| rob_model: Loaded RoBERTa model. | |
| rob_tok: RoBERTa tokenizer. | |
| Returns: | |
| Dict with keys: | |
| - "probs": np.ndarray of shape (5,) | |
| - "top1": str, e.g. "B" | |
| - "top3": list of str, e.g. ["B", "A", "C"] | |
| - "confidence": float, top-1 probability | |
| """ | |
| probs = get_probs(rob_model, rob_tok, prompt, options) | |
| top3_idx = np.argsort(-probs)[:3] | |
| return { | |
| "probs": probs, | |
| "top1": ID2LABEL[int(top3_idx[0])], | |
| "top3": [ID2LABEL[int(i)] for i in top3_idx], | |
| "confidence": float(probs[top3_idx[0]]), | |
| } | |
| def top3_string(probs: np.ndarray) -> str: | |
| """Return space-separated top-3 predicted option letters (e.g. 'B A C').""" | |
| order = np.argsort(-probs)[:3] | |
| return " ".join(ID2LABEL[i] for i in order) | |