#!/usr/bin/env python3 """Offline Qwen2.5-14B-AWQ + book-RAG submission for the IOL-AI challenge.""" from __future__ import annotations import argparse import csv import io import json import os import re import sys from dataclasses import dataclass from pathlib import Path from typing import Any # The evaluation container has no internet access. Fail locally instead of waiting # for network retries if a model/tokenizer file was not included in the repository. os.environ.setdefault("HF_HUB_OFFLINE", "1") os.environ.setdefault("TRANSFORMERS_OFFLINE", "1") os.environ.setdefault("HF_DATASETS_OFFLINE", "1") os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") REPO_DIR = Path(__file__).resolve().parent RESOURCE_DIR = REPO_DIR / "rag_resources" DEFAULT_INPUT = Path(os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv")) DEFAULT_OUTPUT = Path(os.environ.get("IOL_SUBMISSION_CSV", "submission.csv")) from rag_resources.retriever import BookRetriever # noqa: E402 @dataclass class ParseResult: thinking_trace: str final_text: str answers: list[str] valid: bool error: str = "" def _load_config() -> dict[str, Any]: with (RESOURCE_DIR / "config.json").open("r", encoding="utf-8") as handle: config = json.load(handle) integer_overrides = { "IOL_TOP_METHODS": "top_methods", "IOL_TOP_EXAMPLES": "top_examples", "IOL_RAG_MAX_CHARS": "rag_max_chars", "IOL_MAX_NEW_TOKENS": "max_new_tokens", "IOL_EXPLANATION_MAX_NEW_TOKENS": "explanation_max_new_tokens", "IOL_MAX_ATTEMPTS": "max_attempts", "IOL_SEED": "seed", } for environment_name, config_name in integer_overrides.items(): if environment_name in os.environ: config[config_name] = int(os.environ[environment_name]) if "IOL_CHAR_TFIDF_WEIGHT" in os.environ: config["char_tfidf_weight"] = float(os.environ["IOL_CHAR_TFIDF_WEIGHT"]) if "IOL_ENABLE_EXPLANATIONS" in os.environ: config["enable_explanations"] = os.environ[ "IOL_ENABLE_EXPLANATIONS" ].strip().casefold() in {"1", "true", "yes", "on"} return config def _clean_answer_line(line: str) -> str: value = line.strip() value = re.sub(r"^(?:[-*•]\s+)", "", value) value = re.sub(r"^(?:\(?\d+\)?|\(?[A-Za-z]\)?)[.):]\s+", "", value) value = value.strip() pairs = {'"': '"', "'": "'", "`": "`", "“": "”", "‘": "’"} if len(value) >= 2 and value[0] in pairs and value[-1] == pairs[value[0]]: value = value[1:-1].strip() return value def parse_model_response(thinking_trace: str, final_text: str) -> ParseResult: """Parse a separate reasoning trace and a strict FINAL ANSWERS block.""" marker_matches = list( re.finditer(r"(?im)^\s*FINAL\s+ANSWERS\s*:\s*", final_text) ) if not marker_matches: return ParseResult( thinking_trace=thinking_trace.strip(), final_text=final_text.strip(), answers=[], valid=False, error="missing FINAL ANSWERS: marker", ) marker = marker_matches[-1] reasoning_outside_think = final_text[:marker.start()].strip() combined_trace = "\n\n".join( part for part in (thinking_trace.strip(), reasoning_outside_think) if part ) answer_block = final_text[marker.end():].strip() answer_block = re.sub(r"^```(?:text)?\s*", "", answer_block, flags=re.I) answer_block = re.sub(r"\s*```\s*$", "", answer_block) if not answer_block: return ParseResult( thinking_trace=combined_trace, final_text=final_text.strip(), answers=[], valid=False, error="empty FINAL ANSWERS block", ) # Be tolerant if the model emits a JSON list even though bare lines were asked for. answers: list[str] = [] parsed_json_list = False if answer_block.startswith("["): try: decoded = json.loads(answer_block) if isinstance(decoded, list): parsed_json_list = True answers = [str(value).strip() for value in decoded if str(value).strip()] except json.JSONDecodeError: pass if not answers and not parsed_json_list: for line in answer_block.splitlines(): stripped = line.strip() if not stripped or stripped.startswith("```"): continue if re.match(r"(?i)^(?:explanation|reasoning|notes?)\s*:", stripped): break cleaned = _clean_answer_line(stripped) if cleaned: answers.append(cleaned) if not answers: return ParseResult( thinking_trace=combined_trace, final_text=final_text.strip(), answers=[], valid=False, error="no non-empty answers in FINAL ANSWERS block", ) return ParseResult( thinking_trace=combined_trace, final_text=final_text.strip(), answers=answers, valid=True, ) def _best_effort_answers(final_texts: list[str]) -> list[str]: """Keep submission shape valid if all strict retries fail.""" for text in reversed(final_texts): candidates = [] for line in text.splitlines(): cleaned = _clean_answer_line(line) if cleaned and not re.match( r"(?i)^(?:final answers|explanation|reasoning)\s*:?'?", cleaned, ): candidates.append(cleaned) if candidates: return [candidates[-1]] return [""] def _version_tuple(version: str) -> tuple[int, int, int]: numbers = [int(value) for value in re.findall(r"\d+", version)[:3]] return tuple((numbers + [0, 0, 0])[:3]) # type: ignore[return-value] def _keep_last_token_hidden_state(_module: Any, inputs: tuple[Any, ...]) -> tuple[Any, ...] | None: """Avoid materializing full-sequence vocabulary logits during generation.""" if not inputs: return None hidden_states = inputs[0] if hidden_states.ndim == 3 and hidden_states.shape[1] > 1: return (hidden_states[:, -1:, :],) + inputs[1:] return None def load_model_and_tokenizer() -> tuple[Any, Any, Any]: """Load the pre-quantized AWQ model shipped in this repository.""" try: import torch import transformers from transformers import AutoModelForCausalLM, AutoTokenizer except ImportError as exc: raise RuntimeError( "PyTorch, Transformers, Accelerate, AutoAWQ, and safetensors must be " "available in the evaluation image." ) from exc if _version_tuple(transformers.__version__) < (4, 37, 0): raise RuntimeError( f"Qwen2.5 requires transformers>=4.37.0; found {transformers.__version__}." ) model_config_path = REPO_DIR / "config.json" if not model_config_path.exists(): raise RuntimeError( "Qwen2.5-14B-Instruct-AWQ files are missing. Put the complete model and " "tokenizer snapshot in the same repository directory as script.py." ) with model_config_path.open("r", encoding="utf-8") as handle: model_metadata = json.load(handle) quantization_method = str( model_metadata.get("quantization_config", {}).get("quant_method", "") ).casefold() if model_metadata.get("model_type") != "qwen2" or quantization_method != "awq": raise RuntimeError( "This script expects the official Qwen2.5-14B-Instruct-AWQ snapshot " "(model_type=qwen2, quant_method=awq)." ) tokenizer = AutoTokenizer.from_pretrained( str(REPO_DIR), local_files_only=True, trust_remote_code=False ) if tokenizer.pad_token_id is None: tokenizer.pad_token_id = tokenizer.eos_token_id model = AutoModelForCausalLM.from_pretrained( str(REPO_DIR), local_files_only=True, trust_remote_code=False, device_map="auto", torch_dtype=torch.float16, low_cpu_mem_usage=True, ).eval() # Transformers 4.44.1 projects every prompt token through Qwen2's large # vocabulary head and then upcasts all logits to FP32. Generation only uses # the final-token logits, so trim the token dimension immediately before # lm_head to avoid a multi-gigabyte prefill allocation on the T4. model.lm_head.register_forward_pre_hook(_keep_last_token_hidden_state) input_device = model.get_input_embeddings().weight.device return model, tokenizer, input_device def _problem_prompt(row: dict[str, Any], rag_text: str, retry_note: str = "") -> str: work_language = str(row.get("work_lang", "")).strip() task_language = str(row.get("task_lang", "")).strip() language_lines = [] if work_language: language_lines.append(f"Working language: {work_language}") if task_language: language_lines.append(f"Problem language: {task_language}") language_metadata = "\n".join(language_lines) sections = [rag_text] if rag_text else [] sections.append("CURRENT PROBLEM") if language_metadata: sections.append(language_metadata) sections.extend( [ "CONTEXT:\n" + str(row.get("context", "")), "QUERY:\n" + str(row.get("query", "")), ] ) if retry_note: sections.append(retry_note) return "\n\n".join(sections) def _encode_fitted_prompt( tokenizer: Any, system_prompt: str, row: dict[str, Any], rag_text: str, retry_note: str, context_window: int, requested_new_tokens: int, ) -> tuple[dict[str, Any], int]: """Trim retrieved material, never the current problem, to fit the context.""" fitted_rag = rag_text while True: messages = [ {"role": "system", "content": system_prompt}, {"role": "user", "content": _problem_prompt(row, fitted_rag, retry_note)}, ] rendered = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) encoded = tokenizer(rendered, return_tensors="pt", add_special_tokens=False) prompt_tokens = int(encoded["input_ids"].shape[-1]) available = context_window - prompt_tokens if available >= requested_new_tokens: return encoded, requested_new_tokens if fitted_rag: if len(fitted_rag) <= 600: fitted_rag = "" else: new_length = max(0, int(len(fitted_rag) * 0.72)) fitted_rag = fitted_rag[:new_length].rsplit("\n", 1)[0].rstrip() fitted_rag += "\n[retrieved material shortened to fit the model context]" continue if available >= 256: return encoded, available raise RuntimeError( f"The current problem and system prompt use {prompt_tokens} tokens, " f"leaving only {available} tokens in the {context_window}-token context." ) def generate_once( model: Any, tokenizer: Any, input_device: Any, system_prompt: str, row: dict[str, Any], rag_text: str, retry_note: str, config: dict[str, Any], seed: int, ) -> ParseResult: import torch encoded, max_new_tokens = _encode_fitted_prompt( tokenizer=tokenizer, system_prompt=system_prompt, row=row, rag_text=rag_text, retry_note=retry_note, context_window=int(config["model_context_tokens"]), requested_new_tokens=int(config["max_new_tokens"]), ) encoded = {name: value.to(input_device) for name, value in encoded.items()} torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) # These values match the generation configuration shipped with Qwen2.5 # Instruct checkpoints while keeping an explicit challenge-time token cap. with torch.inference_mode(): output = model.generate( **encoded, max_new_tokens=max_new_tokens, do_sample=True, temperature=float(config["temperature"]), top_p=float(config["top_p"]), top_k=int(config["top_k"]), repetition_penalty=float(config["repetition_penalty"]), use_cache=True, pad_token_id=tokenizer.pad_token_id, ) prompt_length = int(encoded["input_ids"].shape[-1]) generated_ids = output[0, prompt_length:].tolist() response_text = tokenizer.decode(generated_ids, skip_special_tokens=True).strip() return parse_model_response("", response_text) def solve_row( row: dict[str, Any], row_index: int, retriever: BookRetriever, model: Any, tokenizer: Any, input_device: Any, system_prompt: str, config: dict[str, Any], ) -> ParseResult: retrieval = retriever.retrieve( row, top_methods=int(config["top_methods"]), top_examples=int(config["top_examples"]), char_tfidf_weight=float(config["char_tfidf_weight"]), ) rag_text = retriever.format_for_prompt( retrieval, max_chars=int(config["rag_max_chars"]) ) failed_results: list[ParseResult] = [] last_error = "" max_attempts = int(config["max_attempts"]) for attempt in range(max_attempts): retry_note = "" if attempt: retry_note = ( "FORMAT RETRY. The previous sample was rejected by the answer parser " f"because it had {last_error}. Generate a fresh solution. Even if " "uncertain, end with exactly FINAL ANSWERS: followed by at least one " "non-empty bare answer line and nothing else in that block." ) result = generate_once( model=model, tokenizer=tokenizer, input_device=input_device, system_prompt=system_prompt, row=row, rag_text=rag_text, retry_note=retry_note, config=config, seed=int(config["seed"]) + row_index * max_attempts + attempt, ) if result.valid: if attempt: print(f" parser accepted retry {attempt + 1}/{max_attempts}", flush=True) return result last_error = result.error failed_results.append(result) print( f" parser rejected attempt {attempt + 1}/{max_attempts}: {last_error}", flush=True, ) print(" WARNING: all strict retries failed; using a best-effort final line", flush=True) last_result = failed_results[-1] return ParseResult( thinking_trace=last_result.thinking_trace, final_text=last_result.final_text, answers=_best_effort_answers( [failed_result.final_text for failed_result in failed_results] ), valid=False, error=f"all {max_attempts} strict parsing attempts failed", ) def generate_explanation( solution: ParseResult, model: Any, tokenizer: Any, input_device: Any, config: dict[str, Any], ) -> str: """Summarize the model's accepted reasoning for the optional jury track.""" import torch source_reasoning = solution.thinking_trace.strip() or solution.final_text.strip() answer_text = "\n".join(solution.answers) messages = [ { "role": "system", "content": ( "Write a short, human-readable explanation of the solution in a few " "concise bullet points. State the linguistic rule or pattern found, " "the key evidence, and how the final answers follow. Do not reproduce " "the raw reasoning trace, trial and error, or meta-commentary. Do not " "change the final answers. Output only the explanation." ), }, { "role": "user", "content": ( f"MODEL REASONING:\n{source_reasoning}\n\n" f"FINAL ANSWERS:\n{answer_text}" ), }, ] rendered = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) encoded = tokenizer(rendered, return_tensors="pt", add_special_tokens=False) encoded = {name: value.to(input_device) for name, value in encoded.items()} prompt_tokens = int(encoded["input_ids"].shape[-1]) available = int(config["model_context_tokens"]) - prompt_tokens max_new_tokens = min(int(config["explanation_max_new_tokens"]), available) if max_new_tokens < 32: return ( "- Inferred the relevant linguistic patterns from the supplied examples.\n" "- Applied those patterns to produce the listed answers." ) with torch.inference_mode(): output = model.generate( **encoded, max_new_tokens=max_new_tokens, do_sample=False, use_cache=True, pad_token_id=tokenizer.pad_token_id, ) generated_ids = output[0, prompt_tokens:].tolist() explanation = tokenizer.decode( generated_ids, skip_special_tokens=True ).strip() if explanation: return explanation return ( "- Inferred the relevant linguistic patterns from the supplied examples.\n" "- Applied those patterns to produce the listed answers." ) def read_test_rows(path: Path) -> list[dict[str, Any]]: with path.open("r", encoding="utf-8-sig", newline="") as handle: return list(csv.DictReader(handle)) def _write_submission_stream(handle: Any, rows: list[dict[str, str]]) -> None: writer = csv.DictWriter(handle, fieldnames=["id", "pred", "explanation"]) writer.writeheader() writer.writerows(rows) def write_submission(path: Path, rows: list[dict[str, str]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary_path = path.with_name(f".{path.name}.tmp") with temporary_path.open("w", encoding="utf-8", newline="") as handle: _write_submission_stream(handle, rows) os.replace(temporary_path, path) def run_self_test() -> None: import numpy as np hidden_states = np.arange(24).reshape(1, 3, 8) trimmed = _keep_last_token_hidden_state(None, (hidden_states,)) assert trimmed is not None and trimmed[0].shape == (1, 1, 8) assert np.array_equal(trimmed[0], hidden_states[:, -1:, :]) assert _keep_last_token_hidden_state(None, (hidden_states[:, -1:, :],)) is None parsed = parse_model_response( "A useful analysis.", "One last check.\nFINAL ANSWERS:\n1. čha\n2) multi word form", ) assert parsed.valid assert parsed.answers == ["čha", "multi word form"] assert "One last check." in parsed.thinking_trace assert not parse_model_response("", "FINAL ANSWERS:\n").valid assert not parse_model_response("", "FINAL ANSWERS:\n[]").valid assert not parse_model_response("", 'FINAL ANSWERS:\n[""]').valid assert not parse_model_response("", "The answer is x.").valid retriever = BookRetriever(RESOURCE_DIR) canary_row = { "context": "The following words are number expressions in an unknown language.", "query": "Determine the rule and write the number 25 in words.", "answer": "PRIVATE_ANSWER_CANARY", } result = retriever.retrieve( canary_row, top_methods=2, top_examples=2, char_tfidf_weight=3.0 ) prompt = retriever.format_for_prompt(result, max_chars=8000) assert result["methods"] and result["examples"] assert "PRIVATE_ANSWER_CANARY" not in prompt assert all(example.get("string_only_usable") for example in result["examples"]) exact_example = retriever.examples[0] exact_result = retriever.retrieve( {"context": exact_example["context"], "query": exact_example["query"]}, top_methods=1, top_examples=1, char_tfidf_weight=3.0, ) assert exact_result["examples"][0]["id"] == exact_example["id"] submission_buffer = io.StringIO(newline="") _write_submission_stream( submission_buffer, [ { "id": "007", "pred": json.dumps(["čha", "multi word"], ensure_ascii=False), "explanation": "- Identified the relevant pattern.", } ], ) submission_buffer.seek(0) assert list(csv.DictReader(submission_buffer)) == [ { "id": "007", "pred": '["čha", "multi word"]', "explanation": "- Identified the relevant pattern.", } ] incremental_rows = [ {"id": "001", "pred": "[]", "explanation": ""}, {"id": "002", "pred": "[]", "explanation": ""}, ] incremental_rows[0]["pred"] = json.dumps(["answer"], ensure_ascii=False) incremental_buffer = io.StringIO(newline="") _write_submission_stream(incremental_buffer, incremental_rows) incremental_buffer.seek(0) assert list(csv.DictReader(incremental_buffer)) == incremental_rows print( "Self-test passed: last-token logits hook, parser, empty-answer rejection, " "hybrid book-only retrieval, exact-match retrieval, and submission CSV " "serialization." ) def main() -> None: if hasattr(sys.stdout, "reconfigure"): sys.stdout.reconfigure(encoding="utf-8") parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input", type=Path, default=DEFAULT_INPUT) parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT) parser.add_argument( "--explanations", choices=("on", "off"), default=None, help="Generate optional jury-track explanations (default: config setting)", ) parser.add_argument( "--self-test", action="store_true", help="Test parser/retrieval without loading Qwen" ) args = parser.parse_args() if args.self_test: run_self_test() return if not args.input.exists(): raise SystemExit(f"Test file not found: {args.input}") config = _load_config() explanations_enabled = ( args.explanations == "on" if args.explanations is not None else bool(config.get("enable_explanations", False)) ) system_prompt = (RESOURCE_DIR / "system_prompt.txt").read_text(encoding="utf-8").strip() test_rows = read_test_rows(args.input) retriever = BookRetriever(RESOURCE_DIR) # Do not create an all-empty submission before the model is known to load. # Otherwise a missing/incompatible shard or startup OOM is silently scored as # zero EM and zero chrF instead of being reported as a runtime failure. model, tokenizer, input_device = load_model_and_tokenizer() submission_rows: list[dict[str, str]] = [ { "id": str(row.get("id", "")), "pred": "[]", "explanation": "", } for row in test_rows ] write_submission(args.output, submission_rows) print( f"Initialized incremental submission with {len(submission_rows)} rows at " f"{args.output}; explanations={'on' if explanations_enabled else 'off'}", flush=True, ) for index, row in enumerate(test_rows): row_id = str(row.get("id", "")) print(f"[{index + 1}/{len(test_rows)}] solving id={row_id}", flush=True) solution = solve_row( row=row, row_index=index, retriever=retriever, model=model, tokenizer=tokenizer, input_device=input_device, system_prompt=system_prompt, config=config, ) explanation = "" if explanations_enabled: explanation = generate_explanation( solution=solution, model=model, tokenizer=tokenizer, input_device=input_device, config=config, ) submission_rows[index] = { "id": row_id, "pred": json.dumps(solution.answers, ensure_ascii=False), "explanation": explanation, } write_submission(args.output, submission_rows) print(f"Checkpointed prediction {index + 1}/{len(test_rows)}", flush=True) print(f"Wrote {len(submission_rows)} predictions to {args.output}", flush=True) if __name__ == "__main__": main()