""" ielts_task2_inference.py — Dashboard inference module for IELTS Task 2 scorer. Load once at server startup, then call scorer.predict(input_dict) per request. Input dict keys: prompt (str) — IELTS Task 2 question shown to the student essay (str) — student essay text Output dict keys: ta_score (float) — Task Achievement band score (3.5-9.0, 0.5 steps) cc_score (float) — Coherence & Cohesion band score lr_score (float) — Lexical Resource band score gra_score (float) — Grammatical Range & Accuracy band score overall_band (float) — Overall IELTS band score ta_feedback (str) — Task Achievement feedback cc_feedback (str) — Coherence & Cohesion feedback lr_feedback (str) — Lexical Resource feedback gra_feedback (str) — Grammatical Range & Accuracy feedback strengths (str) — Key strengths improvements (str) — Areas for improvement raw_output (str) — Full model output string """ import os import re import torch from transformers import AutoTokenizer, AutoModelForCausalLM SYSTEM_PROMPT = """You are an expert IELTS examiner with over 20 years of experience. You must score essays STRICTLY and PRECISELY on the official IELTS 0-9 band scale in 0.5 increments. Do NOT default to average scores. Differentiate clearly between band levels: - Band 3.5-4.5: major weaknesses, very limited vocabulary, frequent serious errors - Band 5.0-5.5: noticeable weaknesses, limited range, frequent errors - Band 6.0: adequate but with clear limitations in all criteria - Band 6.5-7.0: good with some weaknesses, generally effective - Band 7.5-8.0: very good, minor weaknesses only - Band 8.5-9.0: expert level, sophisticated, near-perfect or perfect essays For every essay you must: 1. Score each criterion on the official IELTS 0-9 band scale (0.5 increments). 2. Provide detailed feedback for each criterion. 3. Give an Overall Band Score. 4. List key strengths. 5. List specific improvements the candidate should make. Always be precise, constructive, and consistent with official IELTS marking standards.""" class IELTSTask2Scorer: """Loads the fine-tuned Gemma Task 2 model and scores student essays.""" MAX_SEQ_LENGTH = 1400 MAX_NEW_TOKENS = 1024 def __init__(self, model_dir: str): """Load the merged model from model_dir. Call once at server startup.""" if os.path.exists(model_dir): model_dir = os.path.abspath(model_dir) print(f"Loading IELTS Task 2 model from: {model_dir}") self.tokenizer = AutoTokenizer.from_pretrained(model_dir) self.model = AutoModelForCausalLM.from_pretrained( model_dir, torch_dtype=torch.bfloat16, device_map="auto", ) self.model.eval() print("Model loaded and ready.") def predict(self, input_dict: dict) -> dict: """Score a single Task 2 student essay. Args: input_dict: dict with keys prompt, essay. Returns: dict with criterion scores, overall band, per-criterion feedback, raw_output. """ prompt = input_dict.get("prompt", "") essay = input_dict.get("essay", "") user_content = ( f"### Task Prompt:\n{prompt}\n\n" f"### Student Essay:\n{essay}\n\n" "### Detailed Evaluation:" ) messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_content}, ] text = self.tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) inputs = self.tokenizer( text=text, return_tensors="pt", truncation=True, max_length=self.MAX_SEQ_LENGTH, ).to(self.model.device) with torch.no_grad(): output_ids = self.model.generate( **inputs, max_new_tokens=self.MAX_NEW_TOKENS, temperature=0.3, top_p=0.9, repetition_penalty=1.1, do_sample=True, pad_token_id=self.tokenizer.eos_token_id, ) generated = output_ids[0][inputs["input_ids"].shape[1]:] raw = self.tokenizer.decode(generated, skip_special_tokens=True) return self._parse(raw) def _parse(self, raw: str) -> dict: """Parse criterion scores and feedback sections from raw model output.""" def extract_score(label): patterns = [ rf"\*\*{label}:\s*([0-9](?:\.5)?)\*\*", rf"\*\*{label}\*\*:\s*([0-9](?:\.5)?)", rf"{label}:\s*([0-9](?:\.5)?)", rf"{label}\s*[-–]\s*([0-9](?:\.5)?)", rf"{label}[^\n]{{0,30}}?([0-9](?:\.5)?)\s*(?:band|score|/9)?", ] for p in patterns: m = re.search(p, raw, re.IGNORECASE) if m: try: return max(3.5, min(9.0, round(float(m.group(1)) * 2) / 2)) except: continue return None def extract_section(header): for p in [ header + r"\*\*[\:\s]+(.*?)(?=\*\*[A-Z]|\Z)", header + r"[\:\*\s]+(.*?)(?=\*\*[A-Z]|\Z)", header + r"[^\n]*\n(.*?)(?=\*\*[A-Z]|\Z)", ]: m = re.search(p, raw, re.IGNORECASE | re.DOTALL) if m and m.group(1) and m.group(1).strip(): return m.group(1).strip() return "" ta_score = extract_score("Task Achievement") cc_score = extract_score("Coherence and Cohesion") lr_score = extract_score("Lexical Resource") gra_score = extract_score("Grammatical Range and Accuracy") ov_score = extract_score("Overall Band Score") if ov_score is None: scores = [s for s in [ta_score, cc_score, lr_score, gra_score] if s is not None] if scores: ov_score = max(3.5, min(9.0, round((sum(scores) / len(scores)) * 2) / 2)) return { "ta_score" : ta_score, "cc_score" : cc_score, "lr_score" : lr_score, "gra_score" : gra_score, "overall_band": ov_score, "ta_feedback" : extract_section(r"Task Achievement"), "cc_feedback" : extract_section(r"Coherence and Cohesion"), "lr_feedback" : extract_section(r"Lexical Resource"), "gra_feedback": extract_section(r"Grammatical Range and Accuracy"), "strengths" : extract_section(r"Strengths"), "improvements": extract_section(r"Areas for Improvement|Improvement"), "raw_output" : raw, }