import os import subprocess import sys def _install_bundled_deps() -> None: """Install transformers from bundled wheels (eval sandbox has no PyPI access).""" wheels_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "wheels") if not os.path.isdir(wheels_dir): return subprocess.run( [ sys.executable, "-m", "pip", "install", "-q", "--no-index", f"--find-links={wheels_dir}", "transformers==4.56.2", ], check=True, ) _install_bundled_deps() import re import csv import json import shutil import tempfile import torch from transformers import AutoTokenizer, AutoModelForCausalLM # The repo is the working directory at run time, and there is no network. os.environ["HF_HUB_OFFLINE"] = "1" os.environ["TRANSFORMERS_OFFLINE"] = "1" MODEL_ID = "." MAX_NEW_TOKENS_REASON = 1200 MAX_NEW_TOKENS_ANSWER = 256 MAX_NEW_TOKENS_EXPLAIN = 600 TEMPERATURE = 0.8 TOP_P = 0.95 MAX_ATTEMPTS = 3 MAX_REASONING_CHARS = 3500 # keep answer-pass context focused SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) PROMPTS_DIR = os.path.join(SCRIPT_DIR, "prompts") KNOWN_TASK_TYPES = ( "translation", "fill_blanks", "match_letters", "text_to_num", "num_to_text", ) # Short pass-role contracts; task-specific procedure is loaded into the user turn. SYSTEM_REASON = ( "You solve International Linguistics Olympiad problems using only the given " "CONTEXT and QUERY. This turn is REASONING ONLY — do not write `FINAL ANSWERS:`." ) SYSTEM_ANSWER = ( "You emit only final answers for International Linguistics Olympiad problems. " "Start with a line that says exactly `FINAL ANSWERS:`, then exactly the requested " "number of bare answer lines — no reasoning, numbering, quotes, or glosses." ) SYSTEM_EXPLAIN = ( "You explain International Linguistics Olympiad solutions concisely. " "Output only the explanation — no final answers." ) def load_tokenizer(model_id: str = "."): """Load tokenizer, converting tokenizer.json for older tokenizers if needed.""" tokenizer_path = os.path.join(model_id, "tokenizer.json") with open(tokenizer_path, encoding="utf-8") as handle: data = json.load(handle) merges = data.get("model", {}).get("merges", []) if not merges or not isinstance(merges[0], list): return AutoTokenizer.from_pretrained(model_id) data["model"]["merges"] = [" ".join(piece) for piece in merges] tmpdir = tempfile.mkdtemp() for name in ("tokenizer_config.json", "special_tokens_map.json"): src = os.path.join(model_id, name) if os.path.isfile(src): shutil.copy(src, tmpdir) with open(os.path.join(tmpdir, "tokenizer.json"), "w", encoding="utf-8") as handle: json.dump(data, handle) return AutoTokenizer.from_pretrained(tmpdir) def _prompt_path(pass_name: str, task_type: str) -> str: """Resolve prompts//.txt, falling back to default.txt.""" safe = os.path.basename(task_type.strip()) if pass_name == "explain": path = os.path.join(PROMPTS_DIR, "explain.txt") if os.path.isfile(path): return path raise FileNotFoundError(f"Missing explain prompt at {path}") candidate = os.path.join(PROMPTS_DIR, pass_name, f"{safe}.txt") if safe in KNOWN_TASK_TYPES and os.path.isfile(candidate): return candidate default = os.path.join(PROMPTS_DIR, pass_name, "default.txt") if os.path.isfile(default): return default raise FileNotFoundError( f"No {pass_name} prompt for task_type={task_type!r} under {PROMPTS_DIR}" ) def expected_answer_count(query: str, task_type: str) -> int: if task_type == "match_letters": numbered = re.findall(r"^\s*\d+\.", query, re.MULTILINE) return len(numbered) or 1 if "blanks" in query.lower(): range_match = re.search(r"\((\d+)-(\d+)\)", query) if range_match: return int(range_match.group(2)) - int(range_match.group(1)) + 1 return len(re.findall(r"\(\d+\)", query)) or 1 numbered = re.findall(r"^\s*\d+[.)]", query, re.MULTILINE) return len(numbered) or 1 def load_user_prompt( pass_name: str, context: str, task_type: str, query: str, reasoning: str = "", n_answers: int | None = None, ) -> str: with open(_prompt_path(pass_name, task_type), encoding="utf-8") as handle: template = handle.read() if n_answers is None: n_answers = expected_answer_count(query, task_type) values = { "context": context.strip(), "query": query.strip(), "task_type": task_type.strip(), "reasoning": reasoning.strip(), "n_answers": n_answers, } needed = set(re.findall(r"\{(\w+)\}", template)) return template.format(**{key: values[key] for key in needed}) def truncate_reasoning(reasoning: str, max_chars: int = MAX_REASONING_CHARS) -> str: """Keep the end of reasoning (drafts + conclusions) within a char budget.""" reasoning = reasoning.strip() if len(reasoning) <= max_chars: return reasoning return "…\n" + reasoning[-max_chars:] # Prefer a dedicated header line; also allow same-line answers after the colon. FINAL_ANSWERS_LINE_RE = re.compile( r"(?im)^[^\w\n]*final answers?[^\w\n]*:?[ \t]*(?=\n|$)|" r"(?im)^[^\w\n]*final answers?\s*:\s*" ) FINAL_ANSWERS_INLINE_RE = re.compile( r"(?is)\bfinal answers?\s*:\s*" ) def extract_raw_final(text: str) -> str: """Return text after the last final-answers marker, or '' if none found.""" line_matches = list(FINAL_ANSWERS_LINE_RE.finditer(text)) if line_matches: return text[line_matches[-1].end() :] inline_matches = list(FINAL_ANSWERS_INLINE_RE.finditer(text)) if inline_matches: return text[inline_matches[-1].end() :] return "" def split_single_line_answer(text: str, expected: int, task_type: str) -> list[str]: text = text.strip() if expected <= 1: return [text] def try_split(pattern: str) -> list[str] | None: parts = [part.strip() for part in re.split(pattern, text) if part.strip()] return parts if len(parts) == expected else None if task_type == "match_letters": for pattern in (r"\s+", r",\s*", r";\s*"): if result := try_split(pattern): return result letters = re.findall(r"[A-Za-z]", text) if len(letters) == expected: return [letter.upper() for letter in letters] return [text] if task_type in ("text_to_num", "num_to_text"): for pattern in (r",\s*", r";\s*", r"\s+"): if result := try_split(pattern): return result return [text] for pattern in (r";\s*", r",\s*"): if result := try_split(pattern): return result return [text] def parse_answer_lines(text_after_marker: str, query: str, task_type: str) -> list[str]: """Parse cleaned answer lines from the raw final-answers section.""" answers = [] for line in text_after_marker.splitlines(): stripped_line = line.strip("`").strip() if stripped_line == "": continue match_numbered_prefix = re.match(r"^\s*\d+[.)]\s+(.*)", stripped_line) if match_numbered_prefix: cleaned_line = match_numbered_prefix.group(1).strip() else: cleaned_line = stripped_line cleaned_line = re.sub(r"\*\*", "", cleaned_line).strip() if task_type == "match_letters": parts = [ part.strip("().[]") for part in re.split(r"[\s,;]+", cleaned_line) if part.strip() ] if not ( len(parts) > 1 and all(re.fullmatch(r"[A-Za-z]", part) for part in parts) ): match_letter_word = re.match( r"^\s*(?:\(([A-Za-z])\)|\[([A-Za-z])\]|([A-Za-z]))\.?:?\s*(.*)$", cleaned_line, ) if match_letter_word: letter = ( match_letter_word.group(1) or match_letter_word.group(2) or match_letter_word.group(3) ) cleaned_line = letter.upper() if cleaned_line: answers.append(cleaned_line) expected = expected_answer_count(query, task_type) if len(answers) == 1 and expected > 1: answers = split_single_line_answer(answers[0], expected, task_type) return answers def postprocess_answer(text, query, task_type): """Keep only the content after the last 'FINAL ANSWERS' marker.""" text_after_marker = extract_raw_final(text) if not text_after_marker.strip(): return [] answers = parse_answer_lines(text_after_marker, query, task_type) expected = expected_answer_count(query, task_type) # Clip extras — misalignment from trailing commentary hurts every later slot. if len(answers) > expected: answers = answers[:expected] return answers def generate( tok, model, messages: list[dict], *, max_new_tokens: int, do_sample: bool, ) -> str: ids = tok.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt", ).to(model.device) gen_kwargs = { "max_new_tokens": max_new_tokens, "do_sample": do_sample, } if do_sample: gen_kwargs["temperature"] = TEMPERATURE gen_kwargs["top_p"] = TOP_P with torch.no_grad(): out = model.generate(ids, **gen_kwargs) return tok.decode(out[0][ids.shape[-1] :], skip_special_tokens=True).strip() def generate_reasoning(tok, model, context: str, task_type: str, query: str) -> str: user = load_user_prompt("reason", context, task_type, query) messages = [ {"role": "system", "content": SYSTEM_REASON}, {"role": "user", "content": user}, ] return generate( tok, model, messages, max_new_tokens=MAX_NEW_TOKENS_REASON, do_sample=True, ) def generate_answers( tok, model, context: str, task_type: str, query: str, reasoning: str, ) -> str: n_answers = expected_answer_count(query, task_type) reasoning = truncate_reasoning(reasoning) user = load_user_prompt( "answer", context, task_type, query, reasoning=reasoning, n_answers=n_answers, ) messages = [ {"role": "system", "content": SYSTEM_ANSWER}, {"role": "user", "content": user}, ] # Greedy answer emit: format discipline matters more than diversity here. return generate( tok, model, messages, max_new_tokens=MAX_NEW_TOKENS_ANSWER, do_sample=False, ) def generate_explanation( tok, model, context: str, task_type: str, query: str, reasoning: str, ) -> str: user = load_user_prompt( "explain", context, task_type, query, reasoning=truncate_reasoning(reasoning), ) messages = [ {"role": "system", "content": SYSTEM_EXPLAIN}, {"role": "user", "content": user}, ] return generate( tok, model, messages, max_new_tokens=MAX_NEW_TOKENS_EXPLAIN, do_sample=False, ) tok = load_tokenizer(MODEL_ID) model = AutoModelForCausalLM.from_pretrained( MODEL_ID, torch_dtype=torch.float16, device_map="auto" ).eval() with open("/tmp/data/test.csv", encoding="utf-8", newline="") as f: test_rows = list(csv.DictReader(f)) # Write incrementally so a wall-clock kill still leaves a partial submission.csv. with open("submission.csv", "w", encoding="utf-8", newline="") as f: writer = csv.DictWriter(f, fieldnames=["id", "pred", "explanation"]) writer.writeheader() f.flush() for idx, r in enumerate(test_rows, start=1): task_type = r["task_type"] reason_tmpl = os.path.basename(_prompt_path("reason", task_type)) answer_tmpl = os.path.basename(_prompt_path("answer", task_type)) print( f"{idx}/{len(test_rows)} task={task_type} " f"reason={reason_tmpl} answer={answer_tmpl}", flush=True, ) reasoning = generate_reasoning( tok, model, r["context"], task_type, r["query"], ) answer_text = generate_answers( tok, model, r["context"], task_type, r["query"], reasoning, ) answers = postprocess_answer(answer_text, r["query"], task_type) explanation = generate_explanation( tok, model, r["context"], task_type, r["query"], reasoning, ) print(f" pred={answers!r}", flush=True) writer.writerow( { "id": r["id"], "pred": json.dumps(answers, ensure_ascii=False), "explanation": explanation, } ) f.flush() print("wrote submission.csv", flush=True)