Smart-MCQ-Solver / src /model.py
Aryan-Prajapati004
Migrate to Streamlit app: MCQ solver with RoBERTa+LoRA from Hub
eb04905
Raw
History Blame Contribute Delete
6.75 kB
"""
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)