""" 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 ──────────────────────────────────────────────────────────────── @torch.no_grad() 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)