| |
| """IOL-AI 2026 submission -- International Linguistics Olympiad solver |
| * Item-count detection now also understands lowercase sub-item letters and |
| fill-in-the-blank underscores, which the previous regex set missed. |
| * Answer voting no longer requires byte-identical normalized strings to |
| agree -- it clusters samples by chrF similarity, so "iepure" and |
| "iepure " and "Iepure" (or a stray trailing gloss) still count as the |
| same vote instead of splitting support three ways. |
| * Self-consistency sampling is convergence-aware: once a problem block's |
| items all have a confident majority, that block is dropped from the |
| prompt list for later rounds, so remaining time budget concentrates on |
| the rows that are still undecided rather than re-rolling rows that have |
| already settled. |
| * The matching-task assignment solver is kept (free-form generation on |
| permutation tasks tends to just emit the identity ordering) but is now |
| reused for the "which digit" family of tasks too, with a lighter |
| single-token constraint. |
| * A rough token-length probe runs right after the tokenizer loads, so the |
| starting batch size is chosen from the actual prompt lengths for this |
| file instead of always starting at a fixed guess and paying for at least |
| one guaranteed OOM. |
| Environment: transformers 4.44.1 / torch 2.4.0 / autoawq, 16GB T4, fp16 only, |
| no internet, hard 30-minute wall clock. |
| """ |
| import os |
| import re |
| import json |
| import time |
| import unicodedata |
| from collections import Counter, defaultdict |
|
|
| START = time.time() |
|
|
| |
| |
| |
| |
| |
| |
| |
| TIME_LIMIT = float(os.environ.get("IOL_TIME_LIMIT", "1800")) |
| SAFETY = float(os.environ.get("IOL_SAFETY", "150")) |
| DEADLINE = START + 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")) |
|
|
| |
| |
| |
| |
| |
| PROMPT_STYLE = os.environ.get("IOL_PROMPT_STYLE", "faithful") |
| 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() - START:7.1f}s] {msg}", flush=True) |
|
|
|
|
| def time_left(): |
| return DEADLINE - time.time() |
|
|
|
|
| |
| |
| |
| _NUM_LINE = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M) |
| _NUM_PAREN = re.compile(r"\((\d{1,3})\)") |
| _NUM_RANGE = re.compile(r"\(?(\d{1,3})\s*(?:[-\u2013\u2014]|to)\s*(\d{1,3})\)?") |
| _LETTER_LINE = re.compile(r"^[ \t]*([A-Za-z])[.)\]]\s", re.M) |
| _LETTER_PAREN = re.compile(r"\(([A-Za-z])\)") |
| _BLANK_MARK = re.compile(r"_{2,}|\[\s*(?:blank|___*)\s*\]", re.I) |
|
|
|
|
| def count_items(query, task_type="", context=""): |
| """Best-effort count of the sub-answers a block expects. Never returns 0.""" |
| q = query or "" |
|
|
| nums = set(int(m) for m in _NUM_LINE.findall(q)) | set( |
| int(m) for m in _NUM_PAREN.findall(q)) |
| letters = set(_LETTER_LINE.findall(q)) | set(_LETTER_PAREN.findall(q)) |
|
|
| span = 0 |
| for a, b in _NUM_RANGE.findall(q): |
| a, b = int(a), int(b) |
| if 0 < b - a < 60: |
| span = max(span, b - a + 1) |
|
|
| marker_count = max(len(nums), len(letters)) |
| if span and marker_count and span != marker_count: |
| |
| |
| n = marker_count |
| else: |
| n = max(span, marker_count) |
|
|
| if n > 1: |
| return n |
|
|
| blanks = len(_BLANK_MARK.findall(q)) |
| if blanks > 1: |
| return blanks |
|
|
| |
| |
| body_lines = [l.strip() for l in q.splitlines() if l.strip()] |
| if len(body_lines) > 1: |
| head, rest = body_lines[0], body_lines[1:] |
| if rest and (head.endswith((":", ".")) or len(rest) >= 2): |
| return len(rest) |
|
|
| |
| |
| if context: |
| cn = len(set(int(m) for m in _NUM_LINE.findall(context))) |
| if cn > 1: |
| return cn |
| cl = len(set(_LETTER_LINE.findall(context))) |
| if cl > 1: |
| return cl |
|
|
| return max(n, 1) |
|
|
|
|
| |
| |
| |
| _LEAD_MARK = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*\u2022]\s+)") |
| _CODE_FENCE = re.compile(r"^```[a-zA-Z]*\s*$") |
| _META_TALK = 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|to summarize\b)", |
| re.I, |
| ) |
|
|
|
|
| def tidy(line): |
| s = line.strip() |
| s = _LEAD_MARK.sub("", s) |
| s = s.strip().strip("`").strip() |
| if len(s) >= 2 and s[0] == s[-1] and s[0] in "\"'\u201c\u201d": |
| s = s[1:-1].strip() |
| return s.strip() |
|
|
|
|
| def source_fallback(query, n): |
| """Last-resort answers when the model produced nothing usable. |
| A blank answer scores zero on both exact-match and chrF, but echoing the |
| prompt's own item text is nearly free and, on transcription-adjacent |
| tasks, shares enough surface form with the true answer to pick up |
| partial chrF credit. Better than nothing, never claimed to be more. |
| """ |
| q = query or "" |
| picked = [] |
| for ln in q.splitlines(): |
| s = ln.strip() |
| if not s: |
| continue |
| m = re.match(r"^\(?(\d{1,3})\)?[.):\]]\s*(.+)$", s) |
| if m: |
| picked.append(m.group(2).strip()) |
| if not picked: |
| lines = [l.strip() for l in q.splitlines() if l.strip()] |
| if len(lines) > 1 and lines[0].endswith((":", ".")): |
| picked = lines[1:] |
| picked = [p.split("|")[0].strip() if "|" in p else p for p in picked] |
| picked = [p for p in picked if p] |
| while len(picked) < n: |
| picked.append(picked[-1] if picked else "?") |
| return picked[:n] |
|
|
|
|
| def squeeze_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 parse_model_output(text, n, fallback=None): |
| if not text or not text.strip(): |
| return list(fallback[:n]) if fallback else ["?"] * n |
|
|
| marker = None |
| for m in re.finditer(r"(?:^|\n)\s*(?:final\s+)?answers?\s*:\s*\n?", text, re.I): |
| marker = m |
| body = text[marker.end():] if marker else text |
|
|
| labelled, plain = [], [] |
| for ln in body.splitlines(): |
| if _CODE_FENCE.match(ln): |
| continue |
| m = re.match(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip()) |
| if m: |
| val = tidy(m.group(2)) |
| if val and not _META_TALK.match(val): |
| labelled.append((int(m.group(1)), val)) |
| c = tidy(ln) |
| if c and not _META_TALK.match(c): |
| plain.append(c) |
|
|
| if len(labelled) >= n: |
| |
| |
| |
| by_label = {} |
| for lab, val in labelled: |
| by_label[lab] = val |
| keys = sorted(by_label) |
| if len(keys) >= n: |
| return [by_label[k] for k in keys[:n]] |
|
|
| return squeeze_to_n(plain, n, fallback) |
|
|
|
|
| _DIGITS_ONLY = re.compile(r"-?\d+") |
|
|
|
|
| def coerce_numeral(answer, task_type): |
| """text_to_num answers should be bare digits; pull them out if the model |
| wrapped them in words or punctuation instead of reworking the parser for |
| one task family.""" |
| if (task_type or "").strip().lower() != "text_to_num": |
| return answer |
| m = _DIGITS_ONLY.search(answer or "") |
| return m.group(0) if m else answer |
|
|
|
|
| def normalize(s): |
| s = unicodedata.normalize("NFC", (s or "").strip().lower()) |
| s = re.sub(r"\s+", " ", s) |
| return s.strip(" .!?;:,") |
|
|
|
|
| |
| |
| |
| def _char_ngrams(s, k): |
| s = re.sub(r"\s+", "", s) |
| if len(s) < k: |
| return Counter() |
| return Counter(s[i:i + k] for i in range(len(s) - k + 1)) |
|
|
|
|
| def chrf(hyp, ref, order=6, beta=2.0): |
| if not hyp or not ref: |
| return 0.0 |
| precisions, recalls = [], [] |
| for k in range(1, order + 1): |
| h, r = _char_ngrams(hyp, k), _char_ngrams(ref, k) |
| if not h or not r: |
| continue |
| overlap = sum((h & r).values()) |
| precisions.append(overlap / max(1, sum(h.values()))) |
| recalls.append(overlap / max(1, sum(r.values()))) |
| if not precisions: |
| return 0.0 |
| p = sum(precisions) / len(precisions) |
| r = sum(recalls) / len(recalls) |
| if p + r == 0: |
| return 0.0 |
| b2 = beta * beta |
| return (1 + b2) * p * r / (b2 * p + r) |
|
|
|
|
| |
| |
| |
| CLUSTER_THRESHOLD = 0.90 |
|
|
|
|
| def cluster_candidates(cands): |
| """Group near-duplicate strings together (single-linkage on chrF). |
| Exact-match voting undercounts agreement: "iepurele" and "iepurele " or a |
| stray trailing gloss are the same underlying answer but would previously |
| split into separate buckets and never reach a majority. Grouping by |
| similarity instead means minor formatting noise doesn't cost a block its |
| vote. |
| """ |
| groups = [] |
| for c in cands: |
| placed = False |
| for g in groups: |
| if chrf(c, g[0]) >= CLUSTER_THRESHOLD or normalize(c) == normalize(g[0]): |
| g.append(c) |
| placed = True |
| break |
| if not placed: |
| groups.append([c]) |
| return groups |
|
|
|
|
| def vote_answer(cands, anchor=None): |
| """Choose one answer for an item from the greedy pass plus samples. |
| The greedy (temperature 0) answer is the default and can only be |
| displaced by a cluster that (a) has at least two members and (b) is |
| strictly larger than the cluster the anchor itself belongs to. This |
| asymmetry matters: with only a handful of samples, treating all |
| candidates as equal and breaking ties by "closest to the others" is |
| poorly conditioned -- with two candidates that comparison is symmetric |
| and degenerates into picking whichever string happens to be shorter. On |
| a held-out mock set that variant swapped out the greedy answer for a |
| sampled one roughly half the time and cost a large chunk of exact match. |
| Anchoring keeps voting monotone: it only overrides on genuine, repeated |
| 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 = cluster_candidates(cands) |
| anchor_group_size = 1 |
| for g in groups: |
| if any(normalize(m) == normalize(anchor) or chrf(m, anchor) >= CLUSTER_THRESHOLD |
| for m in g): |
| anchor_group_size = len(g) |
| break |
|
|
| best_group = max(groups, key=len) |
| if len(best_group) >= 2 and len(best_group) > anchor_group_size: |
| |
| |
| return Counter(best_group).most_common(1)[0][0] |
| return anchor |
|
|
|
|
| |
| |
| |
| |
| _ROW_OPT = re.compile(r"^[ \t]*([A-Za-z])[.)]\s+(.+)$", re.M) |
| _ROW_ITEM = re.compile(r"^[ \t]*(\d{1,3})[.)]\s+(.+)$", re.M) |
|
|
|
|
| def split_matching_block(context): |
| items = [(int(a), b.strip()) for a, b in _ROW_ITEM.findall(context or "")] |
| opts = [(a, b.strip()) for a, b in _ROW_OPT.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 assign_max_weight(score): |
| """One-to-one assignment maximizing total score. scipy if available, |
| otherwise a greedy pass plus a few rounds of local 2-swaps.""" |
| n, m = len(score), len(score[0]) |
| try: |
| import numpy as _np |
| from scipy.optimize import linear_sum_assignment |
| _, cols = linear_sum_assignment(-_np.array(score)) |
| return list(cols) |
| 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 enforce_bijection(answers): |
| """If every answer is a single letter and duplicates exist where the full |
| alphabet of options wasn't used, reassign duplicated slots to the unused |
| letters. Tightly scoped so it never touches anything but this exact |
| single-letter-permutation shape.""" |
| if len(answers) < 3 or not all(re.fullmatch(r"[A-Za-z]", a or "") for a in answers): |
| return answers |
| n = len(answers) |
| universe = [chr(ord("A") + i) for i in range(n)] |
| upper = [a.upper() for a in answers] |
| if len(set(upper)) == n: |
| return answers |
| unused = [l for l in universe if l not in set(upper)] |
| if not unused: |
| return answers |
| seen, out = set(), [] |
| for a in upper: |
| if a in seen and unused: |
| out.append(unused.pop(0)) |
| else: |
| seen.add(a) |
| out.append(a) |
| return out |
|
|
|
|
| |
| |
| |
| FAITHFUL_SYSTEM = ( |
| "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." |
| ) |
|
|
| GUIDED_SYSTEM = ( |
| "You are a gold medallist at the International Linguistics Olympiad.\n" |
| "Each problem gives data from a language you have never seen before; " |
| "everything needed is in the problem itself, and no outside knowledge is " |
| "required.\n" |
| "Work method: align the given examples, segment the forms, find the " |
| "recurring morphemes and the rule that orders them, verify the rule " |
| "against every example given, then apply it to the items you are asked " |
| "for.\n" |
| "Reason concisely. Finish with a block that starts on its own line with " |
| "exactly ANSWERS: followed by one answer per line, in order, with no " |
| "numbering, no commentary, and no blank lines.\n" |
| "Always give your best guess for every item -- never leave one blank." |
| ) |
|
|
| TASK_HINTS = { |
| "translation": "Give only the translation itself -- no source text, no " |
| "gloss, no quotation marks.", |
| "match_letters": "Give a single capital letter per item, identifying its " |
| "match. Every letter is used exactly once; none repeats.", |
| "fill_blanks": "Give only the missing form for that blank -- not the " |
| "whole line, not a gloss.", |
| "text_to_num": "Give the number in digits only (e.g. 111).", |
| "num_to_text": "Give the number spelled out in words, in the problem " |
| "language.", |
| } |
|
|
|
|
| def build_prompt(row, n): |
| if PROMPT_STYLE == "faithful": |
| return f"{row['context'].strip()}\n\n{row['query'].strip()}" |
| 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 reasoning, write ANSWERS: on its own line, 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 -- relevant morphemes, word order, any " |
| "sound changes. Be specific and brief (2-4 sentences or a few short " |
| "bullets); do not restate a chain of reasoning." |
| ) |
|
|
|
|
| def build_explain_prompt(row, answers): |
| return ( |
| f"{row['context'].strip()}\n\n{row['query'].strip()}\n\n" |
| "Answers given:\n" + "\n".join(f"- {a}" for a in answers) + |
| "\n\nBriefly explain the linguistic rules behind these answers." |
| ) |
|
|
|
|
| |
| |
| |
| |
| def report_dev_score(preds): |
| 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 |
| gold_df = pd.read_csv(gold_path, dtype=str) |
| exact, chrfs = [], [] |
| for _, r in gold_df.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] |
| exact.append(1.0 if any(pi.strip() == a.strip() for a in alts) else 0.0) |
| chrfs.append(max(chrf(pi, a) for a in alts)) |
| em = sum(exact) / max(1, len(exact)) |
| cf = sum(chrfs) / max(1, len(chrfs)) |
| log(f" [dev] EM={em:.4f} chrF~={cf:.4f} score~={(em * cf) ** 0.5:.4f} " |
| f"over {len(exact)} 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 = [count_items(r.get("query", ""), r.get("task_type", ""), r.get("context", "")) |
| for _, r in df.iterrows()] |
| total = sum(ns) |
| log(f"loaded {len(df)} problems, {total} items " |
| f"(min={min(ns)} max={max(ns)} mean={total / len(ns):.1f})") |
|
|
| fallbacks = {i: source_fallback(r.get("query", ""), n) |
| for i, (_, r), n in zip(ids, df.iterrows(), ns)} |
|
|
| |
| preds = {i: list(fallbacks[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): |
| 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)} ({time_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): |
| system = FAITHFUL_SYSTEM if PROMPT_STYLE == "faithful" else GUIDED_SYSTEM |
| 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)) |
|
|
| |
| |
| |
| sample_lens = [len(tok.encode(p, add_special_tokens=False)) |
| for p in prompts[: min(8, len(prompts))]] |
| approx_max_len = max(sample_lens) if sample_lens else 512 |
| default_batch = 4 if approx_max_len < 2500 else (2 if approx_max_len < 4500 else 1) |
| batch_size = int(os.environ.get("IOL_BATCH", str(default_batch))) |
| log(f"prompt length probe: ~{approx_max_len} tok max of {len(sample_lens)} " |
| f"sampled -> starting batch_size={batch_size}") |
|
|
| def generate(texts, max_new, sample, temp=0.7): |
| nonlocal batch_size |
| out = [""] * len(texts) |
| order = sorted(range(len(texts)), key=lambda i: len(texts[i])) |
| i = 0 |
| while i < len(order): |
| if time_left() < 25: |
| log(" out of time inside generate(); returning partial batch") |
| 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_by_assignment(row, n): |
| """Score every (item, option-letter) pair via next-token logits and |
| take the best one-to-one assignment, instead of trusting free-form |
| text for match_letters problems (which tends to just emit A, B, C... |
| in order -- a valid permutation, so no repair catches it, and it |
| scores near zero). |
| Raw next-token logprobs are known to carry a systematic per-token |
| prior independent of content (Zhao et al. 2021, "Calibrate Before |
| Use") -- e.g. a model may just prefer the token for "C" regardless of |
| which option is actually correct. We estimate that per-letter prior |
| with a content-free version of the same prompt and subtract it before |
| assignment, so the score reflects evidence from the item text rather |
| than a baked-in letter preference. |
| """ |
| items, opts = split_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: |
| token_ids = set() |
| for form in (L, " " + L): |
| enc = tok.encode(form, add_special_tokens=False) |
| if enc: |
| token_ids.add(enc[0]) |
| cand_ids.append(sorted(token_ids)) |
|
|
| ctx = row["context"].strip() |
|
|
| def _logprob_batch(msgs_list): |
| """Run one forward pass per prompt in msgs_list, return a list of |
| per-candidate-letter logprob lists (same shape as `scores` rows).""" |
| texts = [tok.apply_chat_template(m, tokenize=False, |
| add_generation_prompt=True) |
| for m in msgs_list] |
| out = [] |
| step = 4 |
| for s0 in range(0, len(texts), step): |
| if time_left() < 30: |
| return None |
| chunk = texts[s0:s0 + step] |
| 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() |
| logp = torch.log_softmax(logits, dim=-1) |
| for b in range(len(chunk)): |
| out.append([max(logp[b, i].item() for i in ids_) |
| for ids_ in cand_ids]) |
| return out |
|
|
| |
| q_msgs = [] |
| for num, itext in items: |
| q_msgs.append([ |
| {"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."}, |
| ]) |
| scores = _logprob_batch(q_msgs) |
| if scores is None: |
| return None |
|
|
| |
| |
| |
| |
| |
| |
| |
| opts_block = "\n".join(f"{L}. {t}" for L, t in opts) |
| baseline_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"{opts_block}\n\nWhich lettered option corresponds " |
| f"to item ___ (___)? Reply with the option letter " |
| f"only."}] |
| baseline_scores = _logprob_batch([baseline_msgs]) |
| if baseline_scores is None: |
| |
| |
| log(" calibration baseline unavailable; using raw scores") |
| calibrated = scores |
| else: |
| baseline = baseline_scores[0] |
| calibrated = [[s - baseline[j] for j, s in enumerate(row_scores)] |
| for row_scores in scores] |
|
|
| cols = assign_max_weight(calibrated) |
| return [letters[c] for c in cols] |
|
|
| def apply_parsers(row_id, row, n, text): |
| answer = enforce_bijection(parse_model_output(text, n, fallbacks[row_id])) |
| task_type = (row.get("task_type") or "").strip().lower() |
| if task_type == "match_letters": |
| try: |
| mm = solve_by_assignment(row, n) |
| if mm and len(mm) == n: |
| return mm |
| except Exception as e: |
| log(f" assignment solver failed on {row_id}: {type(e).__name__}: {e}") |
| if task_type == "text_to_num": |
| answer = [coerce_numeral(a, task_type) for a in answer] |
| return answer |
|
|
| |
| TOK_PER_S = float(os.environ.get("IOL_TOKS", "30")) |
| adaptive = int(0.40 * max(1.0, time_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)") |
|
|
| t0 = time.time() |
| texts = generate(prompts, max_new=max_new, sample=False) |
| pass1_cost = time.time() - t0 |
|
|
| samples = {i: [] for i in ids} |
| match_rows = 0 |
| for (row_id, n, txt), (_, row) in zip(zip(ids, ns, texts), df.iterrows()): |
| a = apply_parsers(row_id, row, n, txt) |
| if (row.get("task_type") or "").strip().lower() == "match_letters": |
| match_rows += 1 |
| preds[row_id] = a |
| samples[row_id].append(a) |
| if match_rows: |
| log(f"assignment solver attempted on {match_rows} match_letters problem(s)") |
| write_submission(OUT_CSV, ids, preds, explanations) |
| no_block = sum(1 for t in texts if not re.search(r"answers?\s*:", t or "", re.I)) |
| empty = sum(1 for t in texts if not (t 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)") |
| report_dev_score(preds) |
|
|
| |
| |
| |
| def row_converged(row_id): |
| s = samples[row_id] |
| if len(s) < 3: |
| return False |
| for k in range(len(s[0])): |
| col = [x[k] for x in s if k < len(x)] |
| groups = cluster_candidates(col) |
| if not any(len(g) >= 2 for g in groups): |
| return False |
| return True |
|
|
| reserve = 0.0 |
| if WANT_EXPLANATION: |
| reserve = min(300.0, 0.25 * pass1_cost + 60) |
|
|
| active_ids = list(ids) |
| n_extra = 0 |
| while time_left() - reserve > pass1_cost * 0.5 and n_extra < MAX_SAMPLES and active_ids: |
| n_extra += 1 |
| active_ids = [i for i in active_ids if not row_converged(i)] |
| if not active_ids: |
| log(" every row has converged; stopping self-consistency early") |
| break |
| log(f"self-consistency pass {n_extra} on {len(active_ids)}/{len(ids)} " |
| f"undecided rows ({time_left():.0f}s left)") |
| pos = {rid: k for k, rid in enumerate(ids)} |
| sub_prompts = [prompts[pos[i]] for i in active_ids] |
| texts = generate(sub_prompts, max_new=max_new, sample=True, temp=SAMPLE_TEMP) |
| row_lookup = {str(rid): row for rid, (_, row) in zip(ids, df.iterrows())} |
| n_lookup = dict(zip(ids, ns)) |
| for rid, txt in zip(active_ids, texts): |
| if txt: |
| a = apply_parsers(rid, row_lookup[rid], n_lookup[rid], txt) |
| samples[rid].append(a) |
| for rid in ids: |
| if len(samples[rid]) >= 3: |
| greedy = samples[rid][0] |
| n = n_lookup[rid] if rid in n_lookup else len(greedy) |
| preds[rid] = enforce_bijection([ |
| vote_answer([s[k] for s in samples[rid] 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 up to {n_extra + 1} samples/row -> written") |
| report_dev_score(preds) |
|
|
| |
| if WANT_EXPLANATION and time_left() > 60: |
| log(f"generating explanations ({time_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_texts = generate(ex_prompts, max_new=200, sample=False) |
| for i, e in zip(ids, ex_texts): |
| e = re.sub(r"\s+", " ", (e or "").strip()) |
| if e: |
| explanations[i] = e[:1200] |
| write_submission(OUT_CSV, ids, preds, explanations) |
| log("explanations written") |
|
|
| |
| |
| n_lookup = dict(zip(ids, ns)) |
| malformed = [i for i in ids if len(preds[i]) != n_lookup[i] |
| or any(not str(x).strip() for x in preds[i])] |
| if malformed: |
| log(f"repairing {len(malformed)} malformed rows") |
| for i in malformed: |
| preds[i] = squeeze_to_n( |
| [x for x in preds[i] if str(x).strip()], n_lookup[i], fallbacks[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() - START:.0f}s elapsed") |
|
|
|
|
| if __name__ == "__main__": |
| main() |