| |
| """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() |
|
|
| |
| |
| |
| 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")) |
| MAX_SAMPLES = int(os.environ.get("IOL_MAXSAMPLES", "8")) |
| |
| |
| |
| |
| |
| |
| BASELINE_MODE = os.environ.get("IOL_BASELINE", "1") == "1" |
| |
| |
| 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") |
| |
| 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() |
|
|
|
|
| |
| |
| |
| |
|
|
| _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: |
| |
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| |
| 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) |
|
|
|
|
| |
| |
| |
|
|
| _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() |
| |
| |
| 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:] |
| |
| 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 |
|
|
| |
| 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 len(numbered) >= n: |
| by_label = {} |
| for lab, val in numbered: |
| by_label[lab] = val |
| 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: |
| |
| |
| |
| 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(" .!?;:,") |
|
|
|
|
| |
| |
| |
| |
| |
|
|
| 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): |
| """Pick one answer for an item, given the greedy answer plus samples. |
| |
| `anchor` is the greedy (temperature-0) answer and is the default. Sampled |
| answers may only displace it when at least two of them agree on the same |
| normalised form AND that form has strictly more support than the anchor's. |
| |
| This asymmetry is empirically necessary, not decorative. An earlier version |
| treated all candidates equally and fell back to "most central by chrF" when |
| no majority existed. With only a handful of samples that fallback is |
| ill-defined -- with two candidates the pairwise chrF is symmetric, so it |
| degenerated to picking the shorter string -- and it replaced the greedy |
| answer with a temperature-0.7 sample about half the time. Measured on the |
| mock set that cost 4x exact match (EM 0.044 -> 0.011). Anchoring makes |
| voting monotone: it can only fire on genuine agreement. |
| """ |
| cands = [c for c in cands if c and c.strip()] |
| if anchor is None: |
| anchor = cands[0] if cands else "?" |
| if len(cands) < 3: |
| 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) |
|
|
| if best_key is not None and best_n >= 2 and best_n > anchor_support: |
| return Counter(groups[best_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): |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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." |
| ) |
|
|
|
|
| |
| |
| |
| 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." |
| ) |
|
|
|
|
| |
| |
| |
|
|
| 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)} |
|
|
| |
| 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)") |
|
|
| |
| 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" |
|
|
| |
| |
| |
| |
| def _load(dev_map): |
| |
| |
| 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) |
| |
| |
| |
| |
| |
| |
| |
| |
| 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: |
| 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] |
| |
| 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] |
|
|
| |
| |
| |
| |
| |
| |
| |
| 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: |
| |
| |
| 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])) |
| |
| |
| 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) |
| |
| |
| |
| 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) |
|
|
| |
| reserve = 0.0 |
| if WANT_EXPLANATION: |
| reserve = min(300.0, 0.25 * pass1_cost + 60) |
| n_extra = 0 |
| while 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]) >= 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) |
|
|
| |
| 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") |
|
|
| |
| 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() |
|
|