LOL-AI-2026 / script.py
BigRatz's picture
final submission
07cfe5d verified
Raw
History Blame Contribute Delete
11.5 kB
import os
import re
import json
import time
import unicodedata
from collections import Counter, defaultdict
T0 = time.time()
L1 = float(os.environ.get("IOL_TIME_LIMIT", "1800"))
S1 = float(os.environ.get("IOL_SAFETY", "150"))
D1 = T0 + L1 - S1
P1 = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv")
P2 = os.environ.get("IOL_OUT_CSV", "submission.csv")
M1 = os.environ.get("IOL_MODEL", ".")
E1 = os.environ.get("IOL_EXPLAIN", "1") == "1"
X1 = int(os.environ.get("IOL_MAXNEW", "512"))
X2 = int(os.environ.get("IOL_MAXSAMPLES", "24"))
X3 = float(os.environ.get("IOL_TEMP", "0.5"))
X4 = int(os.environ.get("IOL_BATCH", "4"))
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 lg(msg):
print(f"[{time.time() - T0:7.1f}s] {msg}", flush=True)
def lf():
return D1 - time.time()
_LN = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M)
_PN = re.compile(r"\((\d{1,3})\)")
_RG = re.compile(r"\(?(\d{1,3})\s*(?:[-–—]|to)\s*(\d{1,3})\)?")
_LL = re.compile(r"^[ \t]*([A-Z])[.)\]]\s", re.M)
_SP = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*•]\s+)")
_FC = re.compile(r"^```[a-zA-Z]*\s*$")
_CT = 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 d1(q, t="", c=""):
q = q or ""
ln = [int(m) for m in _LN.findall(q)]
pn = [int(m) for m in _PN.findall(q)]
rn = 0
for a, b in _RG.findall(q):
a, b = int(a), int(b)
if 0 < b - a < 60: rn = max(rn, b - a + 1)
cand = max(len(set(ln)), len(set(pn)))
if rn and cand and rn != cand: return cand
cand = max(cand, len(set(_LL.findall(q))))
n = max(rn, cand)
if n > 1: return n
lns = [l.strip() for l in q.splitlines() if l.strip()]
if len(lns) > 1:
h = lns[0]
b = lns[1:] if h.endswith((":", ".")) else lns
if b: return len(b)
if c:
cn = len(set(int(m) for m in _LN.findall(c)))
if cn > 1: return cn
cl = len(set(_LL.findall(c)))
if cl > 1: return cl
return max(n, 1)
def d2(q, n, t=""):
if t.strip().lower() == "match_letters": return ["A"] * n
q = q 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:
lns = [l.strip() for l in q.splitlines() if l.strip()]
if len(lns) > 1 and lns[0].endswith((":", ".")): out = lns[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 d3(s):
s = s.strip()
s = _SP.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 d4(items, n, fb=None):
items = [i for i in items if i and i.strip()]
if len(items) > n: items = items[-n:]
while len(items) < n:
if fb and len(items) < len(fb): items.append(fb[len(items)])
else: items.append(items[-1] if items else "?")
return items[:n]
def d5(text, n, fb=None):
if not text: return list(fb[:n]) if fb 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 _FC.match(ln): continue
mm = re.match(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip())
if mm:
val = d3(mm.group(2))
if val and not _CT.match(val): numbered.append((int(mm.group(1)), val))
c = d3(ln)
if c and not _CT.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 d4(raw, n, fb)
def d6(s):
s = unicodedata.normalize("NFC", (s or "").strip().lower())
s = _SP.sub("", s)
s = re.sub(r"\s+", " ", s)
return s.strip(" .!?;:,")
def d7(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) < 3: return anchor
groups = defaultdict(list)
for c in cands: groups[d6(c)].append(c)
anchor_support = len(groups.get(d6(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 >= 3 and best_n > anchor_support:
return Counter(groups[best_key]).most_common(1)[0][0]
return anchor
def d8(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
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, StoppingCriteria, StoppingCriteriaList
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
df = pd.read_csv(P1, dtype=str).fillna("")
I1 = [str(x) for x in df["id"].tolist()]
N1 = [d1(r.get("query", ""), r.get("task_type", ""), r.get("context", "")) for _, r in df.iterrows()]
total_items = sum(N1)
lg(f"loaded {len(df)} problems, {total_items} items")
S2 = {i: d2(r.get("query", ""), n, r.get("task_type", "")) for i, (_, r), n in zip(I1, df.iterrows(), N1)}
R1 = {i: list(S2[i]) for i in I1}
E2 = {i: "" for i in I1} if E1 else None
d8(P2, I1, R1, E2)
lg(f"wrote placeholder {P2} ({len(I1)} rows)")
lg("loading tokenizer/model ...")
tk = AutoTokenizer.from_pretrained(M1, trust_remote_code=True)
if tk.pad_token is None: tk.pad_token = tk.eos_token
tk.padding_side = "left"
def _ld(dm):
try:
return AutoModelForCausalLM.from_pretrained(M1, torch_dtype=torch.float16, device_map=dm, trust_remote_code=True).eval()
except TypeError:
return AutoModelForCausalLM.from_pretrained(M1, dtype=torch.float16, device_map=dm, trust_remote_code=True).eval()
try:
ml = _ld({"": 0} if torch.cuda.is_available() else "auto")
except Exception as e:
lg(f"pinned load failed ({e}); falling back to auto")
ml = _ld("auto")
lg(f"model ready ({lf():.0f}s left)")
P3 = []
for _, r in df.iterrows():
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()}"}
]
P3.append(tk.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True))
B1 = X4
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
def d9(texts, max_new, sample, temp=0.7):
nonlocal B1
out = [""] * len(texts)
order = sorted(range(len(texts)), key=lambda i: len(texts[i]))
i = 0
while i < len(order):
if lf() < 25: break
idx = order[i:i + B1]
chunk = [texts[j] for j in idx]
try:
enc = tk(chunk, return_tensors="pt", padding=True, truncation=True, max_length=6144).to(ml.device)
kw = dict(max_new_tokens=max_new, pad_token_id=tk.pad_token_id, repetition_penalty=1.0, stopping_criteria=StoppingCriteriaList([Deadline(D1 - 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 = ml.generate(**enc, **kw)
for k, j in enumerate(idx):
out[j] = tk.decode(o[k][enc["input_ids"].shape[1]:], skip_special_tokens=True)
i += B1
except torch.cuda.OutOfMemoryError:
torch.cuda.empty_cache()
if B1 == 1: i += 1
else: B1 = max(1, B1 // 2)
except Exception:
i += B1
return out
lg(f"Pass 1 (greedy) starting... budget: {X1} tokens/item")
t = time.time()
texts = d9(P3, max_new=X1, sample=False)
c1 = time.time() - t
V1 = {i: [] for i in I1}
for i, n, txt in zip(I1, N1, texts):
p1 = [d3(ln) for ln in (txt or "").splitlines() if ln.strip()]
R1[i] = d4(p1, n, S2[i])
V1[i].append(R1[i])
d8(P2, I1, R1, E2)
lg(f"Pass 1 done in {c1:.0f}s. Written to disk.")
reserve = min(300.0, 0.25 * c1 + 60) if E1 else 30.0
n_extra = 0
while lf() - reserve > c1 * 1.25 and n_extra < X2:
n_extra += 1
lg(f"Self-consistency pass {n_extra} starting... ({lf():.0f}s left)")
texts = d9(P3, max_new=X1, sample=True, temp=X3)
for i, n, txt in zip(I1, N1, texts):
if txt:
V1[i].append(d5(txt, n, S2[i]))
for i, n in zip(I1, N1):
if len(V1[i]) >= 3:
greedy = V1[i][0]
voted = [d7([s[k] for s in V1[i] if k < len(s)], anchor=greedy[k] if k < len(greedy) else None) for k in range(n)]
R1[i] = d4(voted, n, S2[i])
d8(P2, I1, R1, E2)
lg(f"Pass {n_extra+1} voted and written.")
if E1 and lf() > 60:
lg(f"Generating explanations ({lf():.0f}s left)...")
ex_sys = "You explain International Linguistics Olympiad solutions to a human judge. State the key rules of the language: morphemes, word order, sound changes. Be concise (2-4 sentences)."
ex_prompts = []
for _, r in df.iterrows():
i = str(r["id"])
msgs = [
{"role": "system", "content": ex_sys},
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}\n\nAnswers given:\n" + "\n".join(f"- {a}" for a in R1[i]) + "\n\nBriefly explain the linguistic rules behind these answers."}
]
ex_prompts.append(tk.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True))
ex_texts = d9(ex_prompts, max_new=200, sample=False)
for i, e in zip(I1, ex_texts):
e = re.sub(r"\s+", " ", (e or "").strip())
if e: E2[i] = e[:1200]
d8(P2, I1, R1, E2)
lg("Explanations written.")
bad = [i for i, n in zip(I1, N1) if len(R1[i]) != n or any(not str(x).strip() for x in R1[i])]
if bad:
lg(f"Repairing {len(bad)} malformed rows")
for i, n in zip(I1, N1):
R1[i] = d4([x for x in R1[i] if str(x).strip()], n, S2[i])
d8(P2, I1, R1, E2)
lg(f"DONE. {len(I1)} rows, {time.time() - T0:.0f}s elapsed.")
if __name__ == "__main__":
main()