Files
iol-2026-qwen14b-entry/script.py

260 lines
10 KiB
Python
Raw Permalink Normal View History

import os
import json
import pandas as pd
import torch
import re
import time
import unicodedata
from collections import Counter, defaultdict
from transformers import AutoTokenizer, AutoModelForCausalLM
# Offline Environment Setup
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
MODEL_ID = "."
# Time Budgets
T0 = time.time()
DEADLINE = T0 + 1650 # 27.5 minutes max limit
EXPLAIN_RESERVE = 240 # Save 4 minutes at the end for Jury Explanations
TEST_CSV = "/tmp/data/test.csv"
OUT_CSV = "submission.csv"
MAX_SAMPLES = 40 # We want maximum possible samples
SAMPLE_TEMP = 0.40 # Sweet spot between 0.35 and 0.50
# 1. Startup Placeholder Creation
if os.path.exists(TEST_CSV):
try:
init_df = pd.read_csv(TEST_CSV, dtype=str).fillna("")
init_rows = [{"id": str(r["id"]), "pred": json.dumps(["?"], ensure_ascii=False), "explanation": "Startup."} for _, r in init_df.iterrows()]
pd.DataFrame(init_rows).to_csv(OUT_CSV, index=False, encoding="utf-8-sig")
except Exception:
pass
print("Loading tokenizer and 14B AWQ model offline...", flush=True)
tok = AutoTokenizer.from_pretrained(MODEL_ID)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.float16,
device_map={"": 0}
).eval()
df = pd.read_csv(TEST_CSV, dtype=str).fillna("")
ids = [str(x) for x in df["id"].tolist()]
# ---- ALIGNMENT & FALLBACK LOGIC ----
_LINE_NUM = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M)
_PAREN_NUM = re.compile(r"\((\d{1,3})\)")
_LINE_LET = re.compile(r"^[ \t]*([A-Z])[.)\]]\s", re.M)
def detect_n_items(query, context=""):
q = query or ""
line = [int(m) for m in _LINE_NUM.findall(q)]
par = [int(m) for m in _PAREN_NUM.findall(q)]
cand = max(len(set(line)), len(set(par)))
cand = max(cand, len(set(_LINE_LET.findall(q))))
n = max(cand, 1)
if n > 1: return n
lines = [l.strip() for l in q.splitlines() if l.strip()]
if len(lines) > 1 and lines[0].endswith((":", ".")):
body = lines[1:]
if body: return len(body)
if context:
cn = len(set(int(m) for m in _LINE_NUM.findall(context)))
if cn > 1: return cn
return max(n, 1)
def extract_item_sources(query, n):
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]
ns = [detect_n_items(r.get("query", ""), r.get("context", "")) for _, r in df.iterrows()]
srcs = {i: extract_item_sources(r.get("query", ""), n) for i, (_, r), n in zip(ids, df.iterrows(), ns)}
# ---- PARSING & CLEANING ----
_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|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 fit_to_n(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 raw_lines(text, n, fb=None):
lines = [clean_line(ln) for ln in (text or "").splitlines() if clean_line(ln)]
return fit_to_n(lines, n, fb)
# ---- VOTING ENGINE ----
def norm(s):
s = unicodedata.normalize("NFC", (s or "").strip().lower())
s = re.sub(r"\s+", " ", s)
return s.strip(" .!?;:,")
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) < 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
def write_submission(preds, expl):
rows = [{"id": i, "pred": json.dumps(preds[i], ensure_ascii=False), "explanation": expl.get(i, "")} for i in ids]
pd.DataFrame(rows).to_csv(OUT_CSV, index=False, encoding="utf-8-sig")
# ---- SYSTEM PROMPTS ----
SYS_TRIVIAL = (
"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."
)
EXPLAIN_SYSTEM = (
"You explain International Linguistics Olympiad solutions to a human judge. "
"State the key rules of the language: morphemes, word order, sound changes. Concise (2-4 sentences)."
)
preds = {i: list(srcs[i]) for i in ids}
expl = {i: "Deducted rules from context." for i in ids}
samples = {i: [] for i in ids}
# ---- PASS 1: GREEDY ANCHOR ----
print("Starting Pass 1 (Greedy Anchor)...", flush=True)
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
if time.time() > DEADLINE: break
msgs = [
{"role": "system", "content": SYS_TRIVIAL},
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}"}
]
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
try:
with torch.no_grad():
out = model.generate(**inputs, max_new_tokens=128, do_sample=False, repetition_penalty=1.0)
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
ans = raw_lines(text, n, srcs[i])
preds[i] = ans
samples[i].append(ans)
except Exception as e:
print(f"Error on {i}: {e}")
write_submission(preds, expl)
print(f"Pass 1 Complete. Time elapsed: {time.time()-T0:.1f}s", flush=True)
# ---- PASSES 2-N: ACCELERATED MULTI-SAMPLE VOTING ----
pass_num = 1
while time.time() < (DEADLINE - EXPLAIN_RESERVE) and pass_num < MAX_SAMPLES:
pass_num += 2 # We are generating 2 samples per pass
print(f"Starting Voting Passes {pass_num-1} & {pass_num} (Time left: {DEADLINE - time.time():.1f}s)...", flush=True)
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
if time.time() > (DEADLINE - EXPLAIN_RESERVE): break
msgs = [
{"role": "system", "content": SYS_TRIVIAL},
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}"}
]
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
try:
# THROUGHPUT MULTIPLIER: num_return_sequences=2 generates 2 samples at once using shared KV cache!
with torch.no_grad():
out = model.generate(
**inputs,
max_new_tokens=128,
do_sample=True,
temperature=SAMPLE_TEMP,
top_p=0.9,
repetition_penalty=1.0,
num_return_sequences=2
)
# Decode and parse both generated sequences
for seq_idx in range(2):
text = tok.decode(out[seq_idx][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
if text:
s_ans = raw_lines(text, n, srcs[i])
samples[i].append(s_ans)
# Apply Voting Immediately
if len(samples[i]) >= 3:
g = samples[i][0]
preds[i] = fit_to_n([vote([s[k] for s in samples[i] if k < len(s)], anchor=g[k] if k < len(g) else None) for k in range(n)], n, srcs[i])
except torch.cuda.OutOfMemoryError:
torch.cuda.empty_cache()
print(f"OOM on {i}, skipping to preserve memory.")
except Exception:
pass
write_submission(preds, expl)
# ---- ISOLATED JURY EXPLANATION PASS ----
if time.time() < DEADLINE:
print(f"Starting Isolated Explanation Pass (Time left: {DEADLINE - time.time():.1f}s)...", flush=True)
for i, n, (_, r) in zip(ids, ns, df.iterrows()):
if time.time() > DEADLINE: break
ans_str = "\n".join(f"- {a}" for a in preds[i])
msgs = [
{"role": "system", "content": EXPLAIN_SYSTEM},
{"role": "user", "content": f"{r['context'].strip()}\n\n{r['query'].strip()}\n\nAnswers given:\n{ans_str}\n\nBriefly explain the linguistic rules."}
]
prompt_text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
inputs = tok(prompt_text, return_tensors="pt", max_length=6144, truncation=True).to(model.device)
try:
with torch.no_grad():
out = model.generate(**inputs, max_new_tokens=150, do_sample=False, repetition_penalty=1.0)
text = tok.decode(out[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
e = re.sub(r"\s+", " ", text)
if e: expl[i] = e[:1000]
except Exception:
pass
write_submission(preds, expl)
print(f"DONE. Completed ~{pass_num} voting passes in {time.time()-T0:.0f}s elapsed.", flush=True)