v14 / script.py
EjZhou's picture
Upload script.py with huggingface_hub
bc19d35 verified
Raw
History Blame Contribute Delete
12.6 kB
"""IOL-AI 2026 — v14: v12 (known-good, 0.1233) + DETERMINISTIC exact-match boosters.
Base = the EXACT v12 pipeline (v7 prompt + count-fix + CSV-safe single-line explanation,
single greedy decode) that scored 0.1233 / EM 0.0609 / explanation_rate 100. v14 adds ONLY
post-processing that targets exact_match, scoped to the two task types you named — the core
prompt that scored 0.1233 is UNCHANGED for every other task:
* text_to_num: model appends a `COMPUTE: e1 | e2` line; we safe-eval it to exact digits
(kills arithmetic slips). The COMPUTE instruction is added ONLY for text_to_num rows.
* match_letters: `repair_bijection` forces a valid letter permutation in the true bijection
case (numbered-context-items == #labels); plus a valid-letters note, ONLY for matching rows.
No global prompt change, no self-consistency — so any delta vs v12 is attributable to the boosters.
"""
import os
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
import re
import csv
import json
import ast as _ast
MODEL_DIR = os.environ.get("IOL_MODEL_DIR", ".")
TEST_CSV = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv")
OUT_CSV = os.environ.get("IOL_OUT_CSV", "submission.csv")
MAX_NEW_TOKENS = int(os.environ.get("IOL_MAX_NEW_TOKENS", "768"))
QUANT = os.environ.get("IOL_QUANT", "4bit")
ANSWER_MARKER = "###ANSWERS###"
# EXACTLY v12/v7's system prompt (scored 0.1233). Do not edit.
SYSTEM_PROMPT = (
"You are an expert competitor at the International Linguistics Olympiad. "
"Each problem gives data from a language you have never seen; deduce its rules "
"using ONLY the data and hints in the problem, then answer EVERY sub-question.\n\n"
"A problem can have MANY sub-questions even when the query is one sentence: e.g. "
"'give the correspondences' expects one answer for EACH numbered item in the data "
"(often a dozen or more). Work out how many answers are required and give exactly "
"that many, one per item, in the order the items appear.\n\n"
"Reason briefly, then end your reply with the answers in EXACTLY this format, with "
"nothing after it:\n"
f"{ANSWER_MARKER}\n"
"1. <answer to item 1>\n"
"2. <answer to item 2>\n"
"(one numbered line per sub-question, in order)\n\n"
"Each answer line holds ONLY the requested form — a word, phrase, number, or "
"letter — with no restating of the question and no commentary. Answer in the "
"language and direction the query asks. For matching items give just the option "
"letter; for number items give digits or the written-out number as asked. Never "
"leave an item blank — always give your best guess."
)
TASK_HINT = {
"translation": "This is a translation task: each answer is only the translated word/phrase.",
"text_to_num": "This is a number task: each answer is only digits (e.g. 42).",
"num_to_text": "This is a number task: each answer is only the number written in the target language's words.",
"match_letters": "This is a matching task: each answer is only the option letter (A, B, C, ...); give one per item in the data.",
"matching": "This is a matching task: each answer is only the option letter; give one per item in the data.",
"fill_blank": "This is a fill-in-the-blank task: each answer is only the missing form.",
"fill_blanks": "This is a fill-in-the-blank task: each answer is only the missing form.",
}
def extract_letter_options(context):
"""Option labels A,B,C... a matching problem offers (contiguous from A), else None."""
found = set()
for line in context.splitlines():
for m in re.finditer(r"(?:^|\s)([A-Z])[.\)]\s+\S", line):
found.add(m.group(1))
if not found:
return None
letters = sorted(found)
if letters != [chr(ord("A") + i) for i in range(len(letters))]:
return None
return letters if 2 <= len(letters) <= 26 else None
def build_messages(row):
"""v12 prompt, with a task-scoped booster instruction ONLY for numbers/matching."""
context = (row.get("context") or "").strip()
query = (row.get("query") or "").strip()
ttype = (row.get("task_type") or "").strip().lower()
system = SYSTEM_PROMPT
hint = TASK_HINT.get(ttype)
if hint:
system = system + "\n\n" + hint
if ttype == "text_to_num":
system += (
"\n\nException to 'nothing after it': after the answer block, add one line:\n"
"COMPUTE: expr1 | expr2 | ...\n"
"each a plain arithmetic expression (digits, + - *, parentheses only) that "
"evaluates to that item's number, one per item, matching the rule you found."
)
if ttype in ("match_letters", "matching"):
opts = extract_letter_options(context)
if opts:
system += (f"\n\nThe ONLY valid answers are these letters: {', '.join(opts)}. "
"Use no other letter.")
return [
{"role": "system", "content": system},
{"role": "user", "content": context + "\n\n" + query},
]
def detect_count(context, query):
q = re.findall(r"(?m)^\s*(\d+)[\.\)]", query)
if q:
return len(q)
par = re.findall(r"\((\d+)\)", query)
if par:
return len(set(par))
c = re.findall(r"(?m)^\s*(\d+)[\.\)]", context)
if c:
return len(c)
return 1
def _clean_answer(s):
s = re.sub(r"^\s*(?:\d+[\.\):]|[-*•])\s*", "", s).strip()
s = re.sub(r"^(?:answer|ans|translation|result)s?\s*[:\-]\s*", "", s, flags=re.I).strip()
return s.strip("\"'“”‘’` ").strip()
def parse_answers(text, min_count=1):
seg = text.rsplit(ANSWER_MARKER, 1)[1] if ANSWER_MARKER in text else text
numbered = {}
for m in re.finditer(r"(?m)^\s*(\d+)[\.\)]\s*(.+?)\s*$", seg):
numbered[int(m.group(1))] = _clean_answer(m.group(2))
if numbered:
answers = [numbered.get(i, "") for i in range(1, max(numbered) + 1)]
else:
lines = [ln.strip() for ln in seg.splitlines() if ln.strip()]
comma_line = next((ln for ln in reversed(lines) if "," in ln), "")
if comma_line:
answers = [_clean_answer(x) for x in comma_line.split(",")]
else:
answers = [_clean_answer(ln) for ln in lines]
answers = [a if a else "?" for a in answers]
if len(answers) < min_count:
answers += ["?"] * (min_count - len(answers))
return answers if answers else ["?"]
def make_explanation(raw, row):
"""CSV-SAFE single-line explanation from the reasoning (before the answer block)."""
head = raw.split(ANSWER_MARKER, 1)[0] if ANSWER_MARKER in raw else raw
head = " ".join(head.split())[:300].strip()
if head:
return head
t = (row.get("task_type") or "linguistic").replace("_", " ")
return f"Inferred the {t} rule from the given examples and applied it to each item."
# ---- deterministic boosters -------------------------------------------------
_ALLOWED_BINOPS = (_ast.Add, _ast.Sub, _ast.Mult)
_NUM_NODE = getattr(_ast, "Num", None) # Python <3.8 (sandbox is 3.10 -> Constant)
def _safe_arithmetic(expr):
try:
tree = _ast.parse(expr.strip(), mode="eval")
except Exception:
return None
def _ev(n):
if isinstance(n, _ast.Expression):
return _ev(n.body)
if isinstance(n, _ast.Constant) and isinstance(n.value, (int, float)):
return n.value
if _NUM_NODE is not None and isinstance(n, _NUM_NODE):
return n.n
if isinstance(n, _ast.BinOp) and isinstance(n.op, _ALLOWED_BINOPS):
l, r = _ev(n.left), _ev(n.right)
if l is None or r is None:
return None
if isinstance(n.op, _ast.Add):
return l + r
if isinstance(n.op, _ast.Sub):
return l - r
return l * r
if isinstance(n, _ast.UnaryOp) and isinstance(n.op, _ast.USub):
v = _ev(n.operand)
return -v if v is not None else None
return None
return _ev(tree)
def apply_compute_overrides(text, answers):
m = re.search(r"(?im)^\s*COMPUTE\s*:\s*(.+)$", text)
if not m:
return answers
exprs = [e.strip() for e in m.group(1).split("|")]
out = list(answers)
for i, e in enumerate(exprs[:len(out)]):
v = _safe_arithmetic(e)
if v is not None and float(v).is_integer():
out[i] = str(int(v))
return out
def repair_bijection(answers, labels):
n = len(labels)
labels_sorted = sorted(labels)
picks = []
for i in range(n):
a = answers[i] if i < len(answers) else ""
f = re.findall(r"[A-Za-z]", a or "")
c = f[0].upper() if f else ""
picks.append(c if c in labels else "")
result = [None] * n
used = set()
for i in range(n):
if picks[i] and picks[i] not in used:
result[i] = picks[i]
used.add(picks[i])
missing = [l for l in labels_sorted if l not in used]
mi = 0
for i in range(n):
if result[i] is None:
result[i] = missing[mi] if mi < len(missing) else labels_sorted[0]
mi += 1
return result
def postprocess(row, answers, raw):
ttype = (row.get("task_type") or "").strip().lower()
context = (row.get("context") or "")
if ttype == "text_to_num":
answers = apply_compute_overrides(raw, answers)
elif ttype in ("match_letters", "matching"):
labels = extract_letter_options(context)
ctx_items = len(re.findall(r"(?m)^\s*\d+\s*[.\)]", context))
if labels and len(labels) >= 2 and ctx_items == len(labels):
answers = repair_bijection(answers, labels)
return answers
def _already_quantized(model_dir):
cfg = os.path.join(model_dir, "config.json")
try:
with open(cfg, encoding="utf-8") as f:
return "quantization_config" in json.load(f)
except Exception:
return False
def load_model():
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
if tok.pad_token_id is None:
tok.pad_token = tok.eos_token
if not torch.cuda.is_available():
return tok, AutoModelForCausalLM.from_pretrained(
MODEL_DIR, torch_dtype=torch.float32).eval()
kwargs = dict(torch_dtype=torch.float16, device_map="auto")
if _already_quantized(MODEL_DIR):
pass
elif QUANT == "4bit":
from transformers import BitsAndBytesConfig
kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True)
return tok, AutoModelForCausalLM.from_pretrained(MODEL_DIR, **kwargs).eval()
def generate_one(tok, model, messages):
import torch
dev = model.device if hasattr(model, "device") else "cpu"
ids = tok.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt").to(dev)
with torch.no_grad():
gen = model.generate(ids, max_new_tokens=MAX_NEW_TOKENS, do_sample=False,
pad_token_id=tok.pad_token_id)
return tok.decode(gen[0][ids.shape[-1]:], skip_special_tokens=True).strip()
def main():
tok, model = load_model()
with open(TEST_CSV, newline="", encoding="utf-8") as f:
rows = list(csv.DictReader(f))
fout = open(OUT_CSV, "w", newline="", encoding="utf-8")
writer = csv.DictWriter(fout, fieldnames=["id", "pred", "explanation"])
writer.writeheader()
fout.flush()
for k, r in enumerate(rows):
context = (r.get("context") or "").strip()
query = (r.get("query") or "").strip()
min_count = detect_count(context, query)
try:
raw = generate_one(tok, model, build_messages(r))
answers = postprocess(r, parse_answers(raw, min_count), raw)
explanation = make_explanation(raw, r)
except Exception as e:
print("row %s fallback: %r" % (r.get("id"), e), flush=True)
answers = ["?"] * min_count
explanation = "Answer derived from the patterns in the examples."
writer.writerow({"id": r["id"],
"pred": json.dumps(answers, ensure_ascii=False),
"explanation": explanation})
fout.flush()
print("%d/%d done" % (k + 1, len(rows)), flush=True)
fout.close()
print("wrote %s (%d rows)" % (OUT_CSV, len(rows)), flush=True)
if __name__ == "__main__":
main()