#!/usr/bin/env python """IOL-AI 2026 submission -- International Linguistics Olympiad solver. Design notes (the eval sandbox is unforgiving, so these matter): * HARD 30-MINUTE LIMIT. A killed process means no score at all, so the script is structured as a monotonically-improving pipeline: it writes a complete, correctly-shaped submission.csv *before* the model is even loaded, then overwrites it after every improvement. Any crash or timeout leaves the best result reached so far on disk. * ALIGNMENT IS EVERYTHING. Each row is a problem block with N numbered items and `pred` must be a JSON list of exactly N answers, in order. One missing line shifts every later answer and zeroes the whole block on both metrics. So N is detected from the query and the model output is force-fitted to it. * NEVER EMIT AN EMPTY STRING. The final score is a geometric mean of exact match and chrF, so an empty answer scores zero on both. A wrong guess is strictly better than a blank. * Environment is transformers 4.44.1 / torch 2.4.0 / autoawq on a 16GB T4 (fp16 only, no bf16, no flash-attn), with no internet. """ import os import re import json import time import unicodedata from collections import Counter, defaultdict T0 = time.time() # The platform allows 30 minutes. Reserve a margin for model load overhead we # can't predict and for the final write; being 60s early costs a little # accuracy, being 1s late costs the entire submission. TIME_LIMIT = float(os.environ.get("IOL_TIME_LIMIT", "1800")) SAFETY = float(os.environ.get("IOL_SAFETY", "150")) DEADLINE = T0 + TIME_LIMIT - SAFETY TEST_CSV = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv") OUT_CSV = os.environ.get("IOL_OUT_CSV", "submission.csv") MODEL_ID = os.environ.get("IOL_MODEL", ".") WANT_EXPLANATION = os.environ.get("IOL_EXPLAIN", "1") == "1" MAX_NEW = int(os.environ.get("IOL_MAXNEW", "900")) # reasoning budget/item MAX_SAMPLES = int(os.environ.get("IOL_MAXSAMPLES", "8")) # self-consistency cap # BASELINE REPLICATION MODE. The organizers' reference script reaches exact match # 0.0729 on the hidden set with THESE EXACT WEIGHTS; our best is 0.0333. Before # adding anything else we need to know whether that number is reproducible by us # at all. This mode replicates their script literally -- trivial prompt, no CoT, # 512 tokens, batch 1 (no padding at all), naive line split, NO forcing to N -- # and changes exactly one thing: repetition_penalty=1.0, our one proven fix. BASELINE_MODE = os.environ.get("IOL_BASELINE", "1") == "1" # v8: ON by default # Lower than the usual 0.7: samples only earn a vote by agreeing with each # other, so keeping them near the greedy mode makes agreement meaningful. SAMPLE_TEMP = float(os.environ.get("IOL_TEMP", "0.5")) os.environ.setdefault("HF_HUB_OFFLINE", "1") os.environ.setdefault("TRANSFORMERS_OFFLINE", "1") os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") # Reduce allocator fragmentation: at batch 4 the T4 has only ~2GB spare. os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") def log(msg): print(f"[{time.time() - T0:7.1f}s] {msg}", flush=True) def left(): return DEADLINE - time.time() # =========================================================================== # Item-count detection (validated: 98.4% of Linguini items land in # correctly-sized blocks) # =========================================================================== _LINE_NUM = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M) _PAREN_NUM = re.compile(r"\((\d{1,3})\)") _RANGE = re.compile(r"\(?(\d{1,3})\s*(?:[-–—]|to)\s*(\d{1,3})\)?") _LINE_LETTER = re.compile(r"^[ \t]*([A-Z])[.)\]]\s", re.M) _PAREN_LETTER = re.compile(r"\(([A-Z])\)") def detect_n_items(query, task_type="", context=""): """How many numbered sub-items this problem asks for. Never < 1.""" q = query or "" line_nums = [int(m) for m in _LINE_NUM.findall(q)] paren_nums = [int(m) for m in _PAREN_NUM.findall(q)] range_n = 0 for a, b in _RANGE.findall(q): a, b = int(a), int(b) if 0 < b - a < 60: range_n = max(range_n, b - a + 1) cand = max(len(set(line_nums)), len(set(paren_nums))) if range_n and cand and range_n != cand: # A stated range ("items 1-4") can disagree with the markers actually # present; the markers are what we have to answer, so they win. return cand cand = max(cand, len(set(_LINE_LETTER.findall(q))), len(set(_PAREN_LETTER.findall(q)))) n = max(range_n, cand) if n > 1: return n # Unnumbered "Translate into X:" followed by one item per line. lines = [l.strip() for l in q.splitlines() if l.strip()] if len(lines) > 1: head = lines[0] body = lines[1:] if head.endswith((":", ".")) else lines if body: return len(body) # Bare instruction ("Determine the correct correspondences."): items are in # the shared context (this is the match_letters shape). if context: c_nums = len(set(int(m) for m in _LINE_NUM.findall(context))) if c_nums > 1: return c_nums c_lets = len(set(_LINE_LETTER.findall(context))) if c_lets > 1: return c_lets return max(n, 1) # =========================================================================== # Output parsing / repair # =========================================================================== _STRIP_PREFIX = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*•]\s+)") _FENCE = re.compile(r"^```[a-zA-Z]*\s*$") _CHATTY = re.compile( r"^\s*(?:here (?:are|is)\b|answers?\s*:?\s*$|explanation\b|note\b|okay\b|" r"solution\b|reasoning\b|analysis\b|translations?\s*:?\s*$|the answers?\b|" r"let me\b|first,|so,|therefore\b|thus\b)", re.I, ) def clean_line(s): s = s.strip() s = _STRIP_PREFIX.sub("", s) s = s.strip().strip("`").strip() if len(s) >= 2 and s[0] == s[-1] and s[0] in "\"'“”": s = s[1:-1].strip() # "word | gloss" answer lines: keep the side being asked for is ambiguous, # so keep the whole line -- chrF still gives partial credit. return s.strip() def extract_item_sources(query, n): """The source text of each numbered item, used as a last-resort fallback. A blank scores zero on both metrics; echoing the item's own source string is strictly better, and on transcription / fill-the-blank tasks the source and the target share a lot of characters, so it collects real chrF credit. """ q = query or "" out = [] for ln in q.splitlines(): s = ln.strip() if not s: continue m = re.match(r"^\(?(\d{1,3})\)?[.):\]]\s*(.+)$", s) if m: out.append(m.group(2).strip()) if not out: lines = [l.strip() for l in q.splitlines() if l.strip()] if len(lines) > 1 and lines[0].endswith((":", ".")): out = lines[1:] # "form | gloss" items: the left side is the thing being asked about. out = [o.split("|")[0].strip() if "|" in o else o for o in out] out = [o for o in out if o] while len(out) < n: out.append(out[-1] if out else "?") return out[:n] def parse_answers(text, n, fallback=None): """Raw model output -> exactly n non-empty answers.""" if not text: return list(fallback[:n]) if fallback else ["?"] * n # Prefer the explicit final block the prompt asks for. m = None for m2 in re.finditer(r"(?:^|\n)\s*(?:final\s+)?answers?\s*:\s*\n?", text, re.I): m = m2 body = text[m.end():] if m else text numbered, raw = [], [] for ln in body.splitlines(): if _FENCE.match(ln): continue mm = re.match(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip()) if mm: val = clean_line(mm.group(2)) if val and not _CHATTY.match(val): numbered.append((int(mm.group(1)), val)) c = clean_line(ln) if c and not _CHATTY.match(c): raw.append(c) # If the model numbered its answers, trust those labels for placement. if len(numbered) >= n: by_label = {} for lab, val in numbered: by_label[lab] = val # last write wins (models restate) labs = sorted(by_label) if len(labs) >= n: return [by_label[l] for l in labs[:n]] return fit_to_n(raw, n, fallback) def fit_to_n(items, n, fallback=None): items = [i for i in items if i and i.strip()] if len(items) > n: # Take the LAST n. The prompt asks for reasoning first and the answers # last, so when there is no ANSWERS: marker to slice on, the tail is the # answer block and the head is reasoning prose. items = items[-n:] while len(items) < n: if fallback and len(items) < len(fallback): items.append(fallback[len(items)]) else: items.append(items[-1] if items else "?") return items[:n] def norm(s): s = unicodedata.normalize("NFC", (s or "").strip().lower()) s = re.sub(r"\s+", " ", s) return s.strip(" .!?;:,") # =========================================================================== # chrF (inline, dependency-free) -- used only to pick the most "central" # candidate when self-consistency voting has no majority. sacrebleu is not # guaranteed to be importable inside the sandbox. # =========================================================================== def _ngrams(s, k): s = re.sub(r"\s+", "", s) return Counter(s[i:i + k] for i in range(len(s) - k + 1)) if len(s) >= k else Counter() def chrf_sim(hyp, ref, order=6, beta=2.0): if not hyp or not ref: return 0.0 ps, rs = [], [] for k in range(1, order + 1): h, r = _ngrams(hyp, k), _ngrams(ref, k) if not h or not r: continue overlap = sum((h & r).values()) ps.append(overlap / max(1, sum(h.values()))) rs.append(overlap / max(1, sum(r.values()))) if not ps: return 0.0 p, r = sum(ps) / len(ps), sum(rs) / len(rs) if p + r == 0: return 0.0 b2 = beta * beta return (1 + b2) * p * r / (b2 * p + r) def vote(cands, anchor=None): cands = [c for c in cands if c and c.strip()] if anchor is None: anchor = cands[0] if cands else "?" if len(cands) < 2: return anchor groups = defaultdict(list) for c in cands: groups[norm(c)].append(c) anchor_support = len(groups.get(norm(anchor), [])) best_key, best_n = None, 0 for k, v in groups.items(): if len(v) > best_n: best_key, best_n = k, len(v) # At pool size 2 (greedy + 1 sample), the only meaningful signal is a # *disagreement*: sample diverges from greedy. We can't get independent # corroboration with a single sample, so require >=2 non-anchor samples # to actually agree with each other before overriding — not just with # the anchor's absence. non_anchor_groups = {k: v for k, v in groups.items() if k != norm(anchor)} if non_anchor_groups: alt_key, alt_n = max(non_anchor_groups.items(), key=lambda kv: len(kv[1])) if alt_n >= 2 and alt_n > anchor_support: return Counter(non_anchor_groups[alt_key]).most_common(1)[0][0] return anchor _OPT_LINE = re.compile(r"^[ \t]*([A-Za-z])[.)]\s+(.+)$", re.M) _ITEM_LINE = re.compile(r"^[ \t]*(\d{1,3})[.)]\s+(.+)$", re.M) def parse_matching_block(context): """For match_letters: the numbered items and the lettered options.""" items = [(int(a), b.strip()) for a, b in _ITEM_LINE.findall(context or "")] opts = [(a, b.strip()) for a, b in _OPT_LINE.findall(context or "")] seen = set() items = [x for x in items if not (x[0] in seen or seen.add(x[0]))] seen = set() opts = [x for x in opts if not (x[0] in seen or seen.add(x[0]))] return items, opts def best_assignment(score): """Max-weight one-to-one assignment. scipy if present, else greedy+swaps.""" n, m = len(score), len(score[0]) try: from scipy.optimize import linear_sum_assignment import numpy as _np r, c = linear_sum_assignment(-_np.array(score)) return list(c) except Exception: pass used, out = set(), [0] * n order = sorted(range(n), key=lambda i: -(max(score[i]) - sorted(score[i])[-2] if m > 1 else 0)) for i in order: j = max((j for j in range(m) if j not in used), key=lambda j: score[i][j], default=0) used.add(j) out[i] = j for _ in range(4): # local 2-swaps improved = False for a in range(n): for b in range(a + 1, n): cur = score[a][out[a]] + score[b][out[b]] alt = score[a][out[b]] + score[b][out[a]] if alt > cur + 1e-9: out[a], out[b] = out[b], out[a] improved = True if not improved: break return out def repair_bijection(answers): """match_letters answers are usually a permutation of the option letters. When every answer is a single letter and there are as many items as distinct letters available, duplicates are certainly wrong. Reassign the duplicated slots to the unused letters. Strictly guarded so it is a no-op on anything that isn't this shape. """ if len(answers) < 3: return answers if not all(re.fullmatch(r"[A-Z]", a or "") for a in answers): return answers n = len(answers) universe = [chr(ord("A") + i) for i in range(n)] if len(set(answers)) == n: return answers unused = [l for l in universe if l not in set(answers)] if not unused: return answers seen, out = set(), [] for a in answers: if a in seen and unused: out.append(unused.pop(0)) else: seen.add(a) out.append(a) return out # =========================================================================== # Prompting # =========================================================================== SYSTEM = ( "You are a gold medallist at the International Linguistics Olympiad.\n" "Each problem gives data from a language you have never seen. Everything " "you need is in the problem itself; no outside knowledge is required or " "allowed.\n" "Method: line up the given examples, segment the words, identify the " "recurring morphemes and the rules that order them, check your rules " "against EVERY example, then apply them to the items asked for.\n" "Be concise while reasoning. Then output a final block that begins with a " "line containing exactly ANSWERS: followed by one answer per line, in the " "order asked, with no numbering, no commentary and no blank lines.\n" "Give your best guess for every item. Never leave one blank." ) # Exact match is half the score, so the answer's *form* matters as much as its # content. test.csv states the task type, so say precisely what a well-formed # answer looks like. Unknown/absent types simply get no hint. TASK_HINTS = { "translation": "Each answer is the translation alone -- no source text, no " "gloss, no notes, no quotation marks.", "match_letters": "Each answer is a single capital letter identifying the " "match for that numbered item. Every letter is used " "exactly once, so no letter may repeat.", "fill_blanks": "Each answer is only the missing form that belongs in that " "blank -- not the whole line, not the gloss.", "text_to_num": "Each answer is written in digits only (e.g. 111).", "num_to_text": "Each answer is the number written out in the problem " "language, words only.", } def build_prompt(row, n): hint = TASK_HINTS.get((row.get("task_type") or "").strip().lower(), "") return ( f"{row['context'].strip()}\n\n{row['query'].strip()}\n\n" f"There are exactly {n} item{'s' if n != 1 else ''} to answer." + (f" {hint}" if hint else "") + f"\nAfter your reasoning, write ANSWERS: on its own line and then exactly " f"{n} line{'s' if n != 1 else ''}, one answer per item, in order." ) EXPLAIN_SYSTEM = ( "You explain International Linguistics Olympiad solutions to a human judge. " "Given a problem and the answers produced, state the key rules of the " "language that justify them: the relevant morphemes, word order and any " "sound changes. Be specific and concise (2-4 sentences or a few short " "bullets). Do not restate the reasoning as a stream of thought." ) def build_explain_prompt(row, answers): return ( f"{row['context'].strip()}\n\n{row['query'].strip()}\n\n" f"Answers given:\n" + "\n".join(f"- {a}" for a in answers) + "\n\nBriefly explain the linguistic rules behind these answers." ) # =========================================================================== # Main # =========================================================================== def dev_score(preds): """Offline diagnostic: score against a gold file when IOL_GOLD is set. Never runs on the platform (the answers are hidden, so the variable is unset there); it exists so one benchmark run reveals the whole learning curve -- greedy, then after each self-consistency pass -- instead of a single final number. """ gold_path = os.environ.get("IOL_GOLD") if not gold_path or not os.path.exists(gold_path): return try: import ast import pandas as pd g = pd.read_csv(gold_path, dtype=str) ems, cfs = [], [] for _, r in g.iterrows(): gold = ast.literal_eval(r["answer"]) p = preds.get(str(r["id"]), []) p = list(p)[:len(gold)] + [""] * max(0, len(gold) - len(p)) for gi, pi in zip(gold, p): alts = gi if isinstance(gi, (list, tuple)) else [gi] alts = [str(a) for a in alts] ems.append(1.0 if any(pi.strip() == a.strip() for a in alts) else 0.0) cfs.append(max(chrf_sim(pi, a) for a in alts)) em = sum(ems) / max(1, len(ems)) cf = sum(cfs) / max(1, len(cfs)) log(f" [dev] EM={em:.4f} chrF~={cf:.4f} score~={(em * cf) ** 0.5:.4f} " f"over {len(ems)} items") except Exception as e: log(f" [dev] scoring failed: {type(e).__name__}: {e}") def write_submission(path, ids, preds, explanations=None): import pandas as pd rows = [] for i in ids: rec = {"id": i, "pred": json.dumps(preds[i], ensure_ascii=False)} if explanations is not None: rec["explanation"] = explanations.get(i, "") rows.append(rec) pd.DataFrame(rows).to_csv(path, index=False) def main(): import pandas as pd df = pd.read_csv(TEST_CSV, dtype=str).fillna("") ids = [str(x) for x in df["id"].tolist()] ns = [detect_n_items(r.get("query", ""), r.get("task_type", ""), r.get("context", "")) for _, r in df.iterrows()] total_items = sum(ns) log(f"loaded {len(df)} problems, {total_items} items " f"(min={min(ns)} max={max(ns)} mean={total_items / len(ns):.1f})") srcs = {i: extract_item_sources(r.get("query", ""), n) for i, (_, r), n in zip(ids, df.iterrows(), ns)} # --- 1. Baseline submission on disk before anything can go wrong -------- preds = {i: list(srcs[i]) for i in ids} explanations = {i: "" for i in ids} if WANT_EXPLANATION else None write_submission(OUT_CSV, ids, preds, explanations) log(f"wrote placeholder {OUT_CSV} ({len(ids)} rows)") # --- 2. Load model ----------------------------------------------------- import torch from transformers import (AutoTokenizer, AutoModelForCausalLM, StoppingCriteria, StoppingCriteriaList) class Deadline(StoppingCriteria): """Abort generation on wall-clock, checked every token. Without this the budget is only checked between batches, so a batch started near the limit runs past it and the platform kills the process. """ def __init__(self, stop_at): self.stop_at = stop_at def __call__(self, input_ids, scores, **kw): return time.time() > self.stop_at log("loading tokenizer/model ...") tok = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True) if tok.pad_token is None: tok.pad_token = tok.eos_token tok.padding_side = "left" # Pin every layer to the GPU. device_map="auto" is free to spill layers to # CPU when it thinks VRAM is tight, and a couple of offloaded layers make # generation ~100x slower without any error -- the worst kind of failure # here. Falling back to "auto" only if the explicit placement fails. def _load(dev_map): # transformers 4.44 (the sandbox) wants torch_dtype=; 5.x renamed it to # dtype=. Accept either so the same file runs in both. try: return AutoModelForCausalLM.from_pretrained( MODEL_ID, torch_dtype=torch.float16, device_map=dev_map, trust_remote_code=True).eval() except TypeError: return AutoModelForCausalLM.from_pretrained( MODEL_ID, dtype=torch.float16, device_map=dev_map, trust_remote_code=True).eval() try: model = _load({"": 0} if torch.cuda.is_available() else "auto") except Exception as e: log(f"pinned load failed ({type(e).__name__}: {e}); falling back to auto") model = _load("auto") devs = set(str(p.device) for p in model.parameters()) log(f"model ready on {sorted(devs)} ({left():.0f}s of budget left)") if any(d.startswith("cpu") or d == "meta" for d in devs): log("WARNING: part of the model is off-GPU; generation will be very slow") if torch.cuda.is_available(): log(f" VRAM allocated {torch.cuda.memory_allocated()/1e9:.2f} GB / " f"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB") prompts = [] for (_, r), n in zip(df.iterrows(), ns): if BASELINE_MODE: msgs = [{"role": "system", "content": "You solve International Linguistics Olympiad problems. " "Answer every numbered item. Put each answer on its own line, " "in order, with no numbering and no extra text."}, {"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}"}] else: msgs = [{"role": "system", "content": SYSTEM}, {"role": "user", "content": build_prompt(r, n)}] prompts.append(tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)) batch_size = 1 if BASELINE_MODE else int(os.environ.get("IOL_BATCH", "4")) def generate(texts, max_new, sample, temp=0.7): """Batched generation with OOM backoff. Returns list of strings.""" nonlocal batch_size out = [""] * len(texts) order = sorted(range(len(texts)), key=lambda i: len(texts[i])) i = 0 while i < len(order): if left() < 25: log(" out of time inside generate(); returning partial") break idx = order[i:i + batch_size] chunk = [texts[j] for j in idx] try: enc = tok(chunk, return_tensors="pt", padding=True, truncation=True, max_length=6144).to(model.device) # repetition_penalty=1.0 EXPLICITLY. Qwen2.5-14B-Instruct-AWQ # ships generation_config.json with repetition_penalty=1.05, # and unlike temperature/top_p/top_k (which greedy ignores, and # which transformers warns about) a repetition penalty IS # applied under greedy decoding -- silently, with no warning. # 34% of the public gold answers repeat a letter 3+ times # (agglutinative morphology like 'ɨmpʼuhurʼu'), so a 5% penalty # pushes the model off exactly the strings we need. kw = dict(max_new_tokens=max_new, pad_token_id=tok.pad_token_id, repetition_penalty=1.0, stopping_criteria=StoppingCriteriaList( [Deadline(DEADLINE - 10)])) if sample: kw.update(do_sample=True, temperature=temp, top_p=0.95) else: kw.update(do_sample=False) with torch.no_grad(): o = model.generate(**enc, **kw) for k, j in enumerate(idx): out[j] = tok.decode(o[k][enc["input_ids"].shape[1]:], skip_special_tokens=True) i += batch_size except torch.cuda.OutOfMemoryError: torch.cuda.empty_cache() if batch_size == 1: log(" OOM at batch=1; skipping this item") i += 1 else: batch_size = max(1, batch_size // 2) log(f" OOM -> batch_size={batch_size}") except Exception as e: # never die mid-run log(f" generate error: {type(e).__name__}: {e}") i += batch_size return out def solve_matching(row, n): """Score every (item, option) pair and take the best one-to-one assignment. Free-form generation fails badly here: measured on the benchmark the model just emits the option labels in order (A, B, C, ... == the identity permutation), which is a *valid* permutation so no repair fires, and it scores ~0. Asking for one letter at a time and reading the next-token distribution turns the task into an assignment problem the model is actually good at, and the one-to-one constraint is then enforced exactly rather than hoped for. """ items, opts = parse_matching_block(row.get("context", "")) if len(items) < 3 or len(opts) < 3 or len(items) != n: return None letters = [o[0] for o in opts] # token id for each option letter, bare and space-prefixed cand_ids = [] for L in letters: ids = set() for form in (L, " " + L): t = tok.encode(form, add_special_tokens=False) if t: ids.add(t[0]) cand_ids.append(sorted(ids)) ctx = row["context"].strip() prompts_m = [] for num, itext in items: msgs = [ {"role": "system", "content": "You match items to their correct counterparts in a " "linguistics problem. Reply with one option letter only."}, {"role": "user", "content": f"{ctx}\n\nWhich lettered option corresponds to item {num} " f"({itext})? Reply with the option letter only."}, ] prompts_m.append(tok.apply_chat_template( msgs, tokenize=False, add_generation_prompt=True)) score = [] bs = 4 for s0 in range(0, len(prompts_m), bs): if left() < 30: return None chunk = prompts_m[s0:s0 + bs] enc = tok(chunk, return_tensors="pt", padding=True, truncation=True, max_length=6144).to(model.device) with torch.no_grad(): logits = model(**enc).logits[:, -1, :].float() logprobs = torch.log_softmax(logits, dim=-1) for b in range(len(chunk)): score.append([max(logprobs[b, i].item() for i in ids) for ids in cand_ids]) col = best_assignment(score) return [letters[c] for c in col] # --- 3. Pass 1: greedy, guarantees a full answer set -------------------- # Size the reasoning budget to the actual problem count. Measured on the # eval hardware (T4, 14B AWQ, batch 4) throughput is ~32 tok/s, so the whole # 30 minutes buys only ~50k generated tokens. With ~16 problem blocks that # affords full-length reasoning; if the platform instead ships one row per # sub-question (~90 rows) a fixed 900-token budget would not even finish a # single pass. Spend at most ~40% of what's left on pass 1. TOK_PER_S = float(os.environ.get("IOL_TOKS", "30")) adaptive = int(0.40 * max(1.0, left()) * TOK_PER_S / max(1, len(df))) max_new = max(192, min(MAX_NEW, adaptive)) log(f"reasoning budget: {max_new} new tokens/problem " f"(adaptive={adaptive}, cap={MAX_NEW}, {len(df)} problems)") t = time.time() texts = generate(prompts, max_new=max_new, sample=False) pass1_cost = time.time() - t samples = {i: [] for i in ids} n_matched = 0 for (i, n, txt), (_, row) in zip(zip(ids, ns, texts), df.iterrows()): if BASELINE_MODE: # literally the organizers' parse: every non-empty stripped line, # however many there are. No cleaning, no fallback, no forcing. preds[i] = [ln.strip() for ln in (txt or "").splitlines() if ln.strip()] samples[i].append(preds[i]) continue a = repair_bijection(parse_answers(txt, n, srcs[i])) # match_letters: free-form generation emits the identity permutation # (A, B, C, ...) and scores ~0, so solve it as an assignment instead. if (row.get("task_type") or "").strip().lower() == "match_letters": try: mm_ = solve_matching(row, n) if mm_ and len(mm_) == n: a = mm_ n_matched += 1 except Exception as e: log(f" matching solver failed on {i}: {type(e).__name__}: {e}") preds[i] = a samples[i].append(a) if n_matched: log(f"assignment solver used on {n_matched} match_letters problem(s)") write_submission(OUT_CSV, ids, preds, explanations) # How often did reasoning run past the token budget before the model got to # its ANSWERS: block? Those problems fall back to salvaged lines, so a high # count means max_new is too small rather than the model being wrong. no_block = sum(1 for txt in texts if not re.search(r"answers?\s*:", txt or "", re.I)) empty = sum(1 for txt in texts if not (txt or "").strip()) log(f"pass 1 (greedy) done in {pass1_cost:.0f}s -> submission written " f"({no_block}/{len(texts)} without an ANSWERS: block, {empty} empty)") dev_score(preds) # --- 4. Self-consistency passes while budget allows --------------------- reserve = 0.0 if WANT_EXPLANATION: reserve = min(300.0, 0.25 * pass1_cost + 60) # explanations are short n_extra = 0 while (not BASELINE_MODE and left() - reserve > pass1_cost * 1.25 and n_extra < MAX_SAMPLES): n_extra += 1 log(f"self-consistency pass {n_extra} ({left():.0f}s left)") texts = generate(prompts, max_new=max_new, sample=True, temp=SAMPLE_TEMP) for i, n, txt in zip(ids, ns, texts): if txt: samples[i].append(repair_bijection(parse_answers(txt, n, srcs[i]))) for i, n in zip(ids, ns): if len(samples[i]) >= 2: # was >= 3 greedy = samples[i][0] preds[i] = repair_bijection( [vote([s[k] for s in samples[i] if k < len(s)], anchor=greedy[k] if k < len(greedy) else None) for k in range(n)]) write_submission(OUT_CSV, ids, preds, explanations) log(f" voted over {n_extra + 1} samples (greedy-anchored) -> written") dev_score(preds) # --- 5. Explanations for the jury track --------------------------------- if WANT_EXPLANATION and left() > 60: log(f"generating explanations ({left():.0f}s left)") ex_prompts = [] for (_, r), i in zip(df.iterrows(), ids): msgs = [{"role": "system", "content": EXPLAIN_SYSTEM}, {"role": "user", "content": build_explain_prompt(r, preds[i])}] ex_prompts.append(tok.apply_chat_template( msgs, tokenize=False, add_generation_prompt=True)) ex = generate(ex_prompts, max_new=200, sample=False) for i, e in zip(ids, ex): e = re.sub(r"\s+", " ", (e or "").strip()) if e: explanations[i] = e[:1200] write_submission(OUT_CSV, ids, preds, explanations) log("explanations written") # --- 6. Final integrity check ------------------------------------------ bad = [i for i, n in zip(ids, ns) if len(preds[i]) != n or any( not str(x).strip() for x in preds[i])] if bad: log(f"repairing {len(bad)} malformed rows") for i, n in zip(ids, ns): preds[i] = fit_to_n([x for x in preds[i] if str(x).strip()], n, srcs[i]) write_submission(OUT_CSV, ids, preds, explanations) log(f"DONE. {len(ids)} rows, {sum(len(v) for v in preds.values())} answers, " f"{time.time() - T0:.0f}s elapsed") if __name__ == "__main__": main()