""" IOL-AI Challenge 2026 — submission script (OFFLINE / Mode B). Runtime facts (Space Submission tab): * T4 medium, 16 GB VRAM, Python 3.10, 30-min wall clock. * NO internet: cannot pip install or download anything. Model weights must be committed into THIS repo (the working dir) and loaded from ".". Only the pre-installed libraries/versions are available (torch 2.4.0, transformers 4.44.1, accelerate 0.34.2, bitsandbytes 0.43.3, autoawq 0.2.7, pandas 2.2.2, numpy 2.1.3, ...). Do NOT pin different majors of torch/transformers/numpy. * Read hidden test set from /tmp/data/test.csv; write submission.csv here. * pred = JSON list, one entry per numbered item, in query order. Ship the model in the repo with build_repo.py. This script loads it from "." with bitsandbytes 4-bit by default (or auto-detected AWQ) so it fits 16 GB. T4 has no bf16 -> use float16. v4: the model's answer COUNT comes from the model, not a fragile query regex — the v1/v2 bug clipped matching/fill-in-blank problems (whose query isn't numbered) down to one answer, zeroing most items. Also: short task_type reminders; single-sequence decode (v3's batched decode OOM'd the T4 -> empty submission = 0); per-row try/except so no single row can zero the whole run; incremental writes + a time-budget guard for the 30-min wall; OPTIONAL sequential self-consistency (IOL_SAMPLES>1, majority vote). Tunable via IOL_SAMPLES / IOL_TEMPERATURE / IOL_TOP_P / IOL_MAX_NEW_TOKENS / IOL_TIME_BUDGET_S. Local dev: set IOL_TEST_CSV to a mock file. Quantization auto-disables if there's no CUDA so the plumbing can be exercised on CPU with a tiny model. """ import os os.environ.setdefault("HF_HUB_OFFLINE", "1") os.environ.setdefault("TRANSFORMERS_OFFLINE", "1") import re import csv import json MODEL_DIR = os.environ.get("IOL_MODEL_DIR", ".") # weights live in the repo TEST_CSV = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv") OUT_CSV = os.environ.get("IOL_OUT_CSV", "submission.csv") MAX_NEW_TOKENS = int(os.environ.get("IOL_MAX_NEW_TOKENS", "768")) # "4bit" (bitsandbytes), "awq" (weights already AWQ-quantized), or "fp16". QUANT = os.environ.get("IOL_QUANT", "4bit") # --- self-consistency knob -------------------------------------------------- # Self-consistency: >1 draws that many SEQUENTIAL sampled decodes per problem # (batch stays 1 -> same VRAM as a single decode, NO OOM risk) and majority-votes # per item. Default 1 = the single greedy decode proven to work in v1/v2. Raise to # 3 only once the Space logs confirm the run finishes comfortably inside 30 min. # (v3 tried BATCHED multi-sequence decode and OOM'd the T4 -> empty submission = 0.) SAMPLES = int(os.environ.get("IOL_SAMPLES", "1")) TEMPERATURE = float(os.environ.get("IOL_TEMPERATURE", "0.7")) TOP_P = float(os.environ.get("IOL_TOP_P", "0.9")) # Safety valve for the 30-min wall: once this many seconds have elapsed, finish # remaining rows with ONE greedy decode instead of SAMPLES sampled ones. TIME_BUDGET_S = float(os.environ.get("IOL_TIME_BUDGET_S", "1620")) # 27 min ANSWER_MARKER = "###ANSWERS###" SYSTEM_PROMPT = ( "You are an expert competitor at the International Linguistics Olympiad. " "Each problem gives data from a language you have never seen; deduce its rules " "using ONLY the data and hints in the problem, then answer EVERY sub-question.\n\n" "A problem can have MANY sub-questions even when the query is one sentence: e.g. " "'give the correspondences' expects one answer for EACH numbered item in the data " "(often a dozen or more). Work out how many answers are required and give exactly " "that many, one per item, in the order the items appear.\n\n" "Reason briefly, then end your reply with the answers in EXACTLY this format, with " "nothing after it:\n" f"{ANSWER_MARKER}\n" "1. \n" "2. \n" "(one numbered line per sub-question, in order)\n\n" "Each answer line holds ONLY the requested form — a word, phrase, number, or " "letter — with no restating of the question and no commentary. Answer in the " "language and direction the query asks. For matching items give just the option " "letter; for number items give digits or the written-out number as asked. Never " "leave an item blank — always give your best guess." ) # Short, low-cost per-task output reminders (the CSV tags each row with task_type). TASK_HINT = { "translation": "This is a translation task: each answer is only the translated word/phrase.", "text_to_num": "This is a number task: each answer is only digits (e.g. 42).", "num_to_text": "This is a number task: each answer is only the number written in the target language's words.", "match_letters": "This is a matching task: each answer is only the option letter (A, B, C, ...); give one per item in the data.", "matching": "This is a matching task: each answer is only the option letter; give one per item in the data.", "fill_blank": "This is a fill-in-the-blank task: each answer is only the missing form.", "fill_blanks": "This is a fill-in-the-blank task: each answer is only the missing form.", } def build_messages(row): """Chat messages for one problem, with a short task_type-specific reminder.""" context = (row.get("context") or "").strip() query = (row.get("query") or "").strip() ttype = (row.get("task_type") or "").strip().lower() system = SYSTEM_PROMPT hint = TASK_HINT.get(ttype) if hint: system = system + "\n\n" + hint return [ {"role": "system", "content": system}, {"role": "user", "content": context + "\n\n" + query}, ] def detect_count(context, query): """Best-effort number of sub-questions — used ONLY as a hint and a minimum pad, NEVER to truncate the model's own answer list (under-producing loses items). Queries usually number items ('1.'/'2.') or mark blanks ('(1)','(2)'); matching queries number nothing, so fall back to the numbered items in the CONTEXT.""" q = re.findall(r"(?m)^\s*(\d+)[\.\)]", query) if q: return len(q) par = re.findall(r"\((\d+)\)", query) if par: return len(set(par)) c = re.findall(r"(?m)^\s*(\d+)[\.\)]", context) if c: return len(c) return 1 def _clean_answer(s): """Strip list markers, common 'Answer:' labels, and surrounding quotes.""" s = re.sub(r"^\s*(?:\d+[\.\):]|[-*•])\s*", "", s).strip() s = re.sub(r"^(?:answer|ans|translation|result)\s*[:\-]\s*", "", s, flags=re.I).strip() return s.strip("\"'“”‘’` ").strip() def parse_answers(text, min_count=1): """Extract the model's FULL answer list — the count comes from the MODEL, never truncated to a query heuristic (that was the v1/v2 bug: it clipped matching problems' dozen answers down to 1). Prefer the ###ANSWERS### block; inside it read the numbered lines; else split a trailing comma-list; else use the lines. Pad up to min_count and never emit a blank.""" seg = text.rsplit(ANSWER_MARKER, 1)[1] if ANSWER_MARKER in text else text numbered = {} for m in re.finditer(r"(?m)^\s*(\d+)[\.\)]\s*(.+?)\s*$", seg): numbered[int(m.group(1))] = _clean_answer(m.group(2)) if numbered: # ordered by the model's indices answers = [numbered.get(i, "") for i in range(1, max(numbered) + 1)] else: lines = [ln.strip() for ln in seg.splitlines() if ln.strip()] comma_line = next((ln for ln in reversed(lines) if "," in ln), "") if comma_line: # matching-style "O, D, A, ..." answers = [_clean_answer(x) for x in comma_line.split(",")] else: answers = [_clean_answer(ln) for ln in lines] answers = [a if a else "?" for a in answers] # never blank (partial credit) if len(answers) < min_count: answers += ["?"] * (min_count - len(answers)) return answers if answers else ["?"] def _norm(s): """Mirror the official scorer's normalization so voting groups answers the same way the metric will (ignore case, surrounding quotes, one trailing dot).""" s = " ".join((s or "").strip().split()) s = s.strip("\"'“”‘’") if s.endswith("."): s = s[:-1] return s.strip().casefold() def vote_answers(sample_texts, min_count=1): """Self-consistency across variable-length answer lists: vote per position on the NORMALIZED form, returning the most common surface form. List length = the longest sample (or min_count). Ties fall to the earliest sample (insertion order).""" from collections import Counter parsed = [parse_answers(t, min_count) for t in sample_texts] n = max([min_count] + [len(p) for p in parsed]) out = [] for i in range(n): counts, surface = Counter(), {} for p in parsed: if i < len(p) and p[i] and p[i] != "?": key = _norm(p[i]) counts[key] += 1 surface.setdefault(key, p[i]) out.append(surface[counts.most_common(1)[0][0]] if counts else "?") return out def _already_quantized(model_dir): """True if the shipped weights are pre-quantized (e.g. AWQ) — then transformers auto-detects the config and we must NOT stack bitsandbytes on top.""" cfg = os.path.join(model_dir, "config.json") try: with open(cfg, encoding="utf-8") as f: return "quantization_config" in json.load(f) except Exception: return False def load_model(): import torch from transformers import AutoTokenizer, AutoModelForCausalLM tok = AutoTokenizer.from_pretrained(MODEL_DIR) if tok.pad_token_id is None: # only used to silence a warning tok.pad_token = tok.eos_token if not torch.cuda.is_available(): model = AutoModelForCausalLM.from_pretrained( MODEL_DIR, torch_dtype=torch.float32).eval() # CPU dev fallback return tok, model kwargs = dict(torch_dtype=torch.float16, device_map="auto") # T4 has no bf16 if _already_quantized(MODEL_DIR): pass # AWQ/pre-quant: transformers reads quantization_config from config.json elif QUANT == "4bit": from transformers import BitsAndBytesConfig kwargs["quantization_config"] = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, ) model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, **kwargs).eval() return tok, model def generate_one(tok, model, messages, do_sample): """Single-sequence decode (batch=1) — the VRAM-safe path proven in v1/v2. (v3's batched multi-sequence decode OOM'd the T4 and produced an empty submission.)""" import torch dev = model.device if hasattr(model, "device") else "cpu" ids = tok.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt").to(dev) gkw = dict(max_new_tokens=MAX_NEW_TOKENS, pad_token_id=tok.pad_token_id) if do_sample: gkw.update(do_sample=True, temperature=TEMPERATURE, top_p=TOP_P) else: gkw.update(do_sample=False) with torch.no_grad(): gen = model.generate(ids, **gkw) return tok.decode(gen[0][ids.shape[-1]:], skip_special_tokens=True).strip() def main(): import time tok, model = load_model() with open(TEST_CSV, newline="", encoding="utf-8") as f: rows = list(csv.DictReader(f)) # Write incrementally so a hard 30-min kill still leaves a valid partial file. fout = open(OUT_CSV, "w", newline="", encoding="utf-8") writer = csv.DictWriter(fout, fieldnames=["id", "pred"]) writer.writeheader() fout.flush() start_t = time.time() for k, r in enumerate(rows): context = (r.get("context") or "").strip() query = (r.get("query") or "").strip() min_count = detect_count(context, query) messages = build_messages(r) # Time guard: once past budget, one greedy decode per remaining row. n_samp = 1 if (time.time() - start_t) > TIME_BUDGET_S else max(1, SAMPLES) try: if n_samp > 1: texts = [generate_one(tok, model, messages, do_sample=True) for _ in range(n_samp)] answers = vote_answers(texts, min_count) else: answers = parse_answers( generate_one(tok, model, messages, do_sample=False), min_count) except Exception as e: # one row must never zero the whole submission print("row %s failed: %r" % (r.get("id"), e), flush=True) answers = ["?"] * min_count writer.writerow({"id": r["id"], "pred": json.dumps(answers, ensure_ascii=False)}) fout.flush() # survive a hard timeout print("%d/%d done" % (k + 1, len(rows)), flush=True) fout.close() print("wrote %s (%d rows)" % (OUT_CSV, len(rows)), flush=True) if __name__ == "__main__": main()