File size: 13,067 Bytes
39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 39a32c4 ed2a871 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 | #!/usr/bin/env python
import os, re, json, time, 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_EXPL = os.environ.get("IOL_EXPLAIN", "1") == "1"
REASON_TOK = int(os.environ.get("IOL_MAXNEW", "900")) # reasoning budget/item
VOTE_CAP = int(os.environ.get("IOL_MAXSAMPLES", "12")) # Phi-4 is fast -> vote a lot
VOTE_TEMP = float(os.environ.get("IOL_TEMP", "0.5"))
BATCH = int(os.environ.get("IOL_BATCH", "4"))
TOK_PER_S = float(os.environ.get("IOL_TOKS", "40")) # Phi-4 throughput est.
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(m): print(f"[{time.time()-T0:7.1f}s] {m}", flush=True)
def left(): return DEADLINE - time.time()
# --- how many items a problem asks for ---
_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)
def count_items(query, ctx=""):
q = query or ""; rng = 0
for a, b in _RG.findall(q):
a, b = int(a), int(b)
if 0 < b - a < 60: rng = max(rng, b - a + 1)
cand = max(len(set(_LN.findall(q))), len(set(_PN.findall(q))), len(set(_LL.findall(q))))
n = max(rng, cand)
if n > 1: return n
lines = [l.strip() for l in q.splitlines() if l.strip()]
if len(lines) > 1:
body = lines[1:] if lines[0].endswith((":", ".")) else lines
if body: return len(body)
if ctx:
cn = len(set(_LN.findall(ctx)))
if cn > 1: return cn
cl = len(set(_LL.findall(ctx)))
if cl > 1: return cl
return max(n, 1)
def item_sources(query, n, task=""):
if (task or "").strip().lower() == "match_letters": return ["A"] * n
q = query or ""; out = []
for ln in q.splitlines():
m = re.match(r"^\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip())
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]
# --- prompt: reason, then a clean ANSWERS: block, shaped by a per-task hint ---
SYS = ("You are a top competitor at the International Linguistics Olympiad. Each problem "
"gives data from a language you have never seen; everything you need is inside the "
"problem. Line up the examples, segment the words, find the recurring morphemes and "
"the rules ordering them, and check them against every example. Reason briefly, then "
"write a line that is exactly ANSWERS: and, below it, one answer per item in the order "
"asked β no numbering, no glosses, no commentary, no blank lines. Never leave one blank.")
HINT = {
"translation": "Each answer is the translated form alone β no source word, no gloss, no quotes.",
"match_letters": "Each answer is one capital letter; every letter is used exactly once, none repeats.",
"fill_blanks": "Each answer is only the missing form for that blank β not the whole line, not the gloss.",
"text_to_num": "Each answer is digits only (for example 111).",
"num_to_text": "Each answer is the number written out in the problem's language, words only.",
}
def user_prompt(row, n):
h = HINT.get((row.get("task_type") or "").strip().lower(), "")
return (f"{str(row['context']).strip()}\n\n{str(row['query']).strip()}\n\n"
f"There are exactly {n} item{'s' if n != 1 else ''} to answer."
+ (f" {h}" if h else "") +
f"\nAfter reasoning, write ANSWERS: on its own line then exactly {n} "
f"line{'s' if n != 1 else ''}, one answer per item, in order.")
# --- parsing the ANSWERS block ---
_PFX = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*β’]\s+)")
_FENCE = re.compile(r"^```")
_PROSE = re.compile(r"^\s*(?:here (?:are|is)\b|answers?\s*:?\s*$|explanation\b|note\b|okay\b|"
r"solution\b|reasoning\b|analysis\b|the answers?\b|let me\b|first,|so,|therefore\b|thus\b)", re.I)
_BRK = re.compile(r"\[[^\[\]\n]*\]")
def clean(s):
s = _PFX.sub("", s.strip()).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 shape(val, task):
t = (task or "").strip().lower()
if t == "match_letters":
m = re.search(r"(?<![A-Za-z])([A-Z])(?![A-Za-z])", val)
return m.group(1) if m else val
if t == "fill_blanks":
m = _BRK.search(val)
if m: return m.group(0)
return val.split("|")[0].strip() if "|" in val else val
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:
items.append(fb[len(items)] if fb and len(items) < len(fb) else (items[-1] if items else "?"))
return items[:n]
def parse(text, n, task, 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 _FENCE.match(ln): continue
mm = re.match(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip())
if mm:
v = clean(mm.group(2))
if v and not _PROSE.match(v): numbered.append((int(mm.group(1)), v))
c = clean(ln)
if c and not _PROSE.match(c): raw.append(c)
if len(numbered) >= n:
by = {}
for lab, v in numbered: by[lab] = v
labs = sorted(by)
if len(labs) >= n: raw = [by[l] for l in labs[:n]]
vals = [shape(v, task) for v in fit_to_n(raw, n, fb)]
return fit_to_n(vals, n, fb)
def norm(s):
s = unicodedata.normalize("NFC", (s or "").strip().lower()); s = re.sub(r"\s+", " ", s)
return s.strip(" .!?;:,")
def vote(cands, anchor):
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
g = defaultdict(list)
for c in cands: g[norm(c)].append(c)
a_sup = len(g.get(norm(anchor), [])); bk, bn = None, 0
for k, v in g.items():
if len(v) > bn: bk, bn = k, len(v)
return Counter(g[bk]).most_common(1)[0][0] if (bk and bn >= 2 and bn > a_sup) else anchor
EXPL_SYS = ("You explain International Linguistics Olympiad solutions to a human judge: state "
"the key rules β morphemes, word order, sound changes β concisely (2-4 sentences). "
"Do not restate the reasoning.")
def write_out(path, ids, preds, expl):
import pandas as pd
rows = [{"id": i, "pred": json.dumps(preds[i], ensure_ascii=False),
"explanation": (expl or {}).get(i, "")} for i in ids]
pd.DataFrame(rows).to_csv(path, index=False)
def main():
import pandas as pd, 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(TEST_CSV, dtype=str).fillna("")
ids = [str(x) for x in df["id"].tolist()]
ns = [count_items(r.get("query", ""), r.get("context", "")) for _, r in df.iterrows()]
tasks = {i: (r.get("task_type", "") or "") for i, (_, r) in zip(ids, df.iterrows())}
srcs = {i: item_sources(r.get("query", ""), n, r.get("task_type", "")) for i, (_, r), n in zip(ids, df.iterrows(), ns)}
log(f"loaded {len(df)} problems, {sum(ns)} items")
preds = {i: list(srcs[i]) for i in ids}; expl = {i: "" for i in ids} if WANT_EXPL else None
write_out(OUT_CSV, ids, preds, expl); log("placeholder written")
class Deadline(StoppingCriteria):
def __init__(self, t): self.t = t
def __call__(self, i, s, **k): return time.time() > self.t
log("loading 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(dm):
try: return AutoModelForCausalLM.from_pretrained(MODEL_ID, torch_dtype=torch.float16, device_map=dm, trust_remote_code=True).eval()
except TypeError: return AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.float16, device_map=dm, 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 ({e}); auto"); model = _load("auto")
log(f"model ready ({left():.0f}s left)")
prompts = [tok.apply_chat_template(
[{"role": "system", "content": SYS}, {"role": "user", "content": user_prompt(r, n)}],
tokenize=False, add_generation_prompt=True) for (_, r), n in zip(df.iterrows(), ns)]
bs = BATCH
def gen(texts, max_new, sample):
nonlocal bs
out = [""] * len(texts); order = sorted(range(len(texts)), key=lambda i: len(texts[i])); i = 0
while i < len(order):
if left() < 25: break
idx = order[i:i+bs]; 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)]))
kw.update(dict(do_sample=True, temperature=VOTE_TEMP, top_p=0.95) if sample else dict(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 += bs
except torch.cuda.OutOfMemoryError:
torch.cuda.empty_cache(); bs = max(1, bs // 2) if bs > 1 else 1
if bs == 1: i += 0
except Exception as e: log(f"gen error: {type(e).__name__}: {e}"); i += bs
return out
adaptive = int(0.40 * max(1.0, left()) * TOK_PER_S / max(1, len(df)))
max_new = max(192, min(REASON_TOK, adaptive))
log(f"Pass 1 greedy, {max_new} tok/item")
t = time.time(); texts = gen(prompts, max_new, sample=False); c1 = time.time() - t
samples = {i: [] for i in ids}
for i, n, txt in zip(ids, ns, texts):
preds[i] = parse(txt, n, tasks[i], srcs[i]); samples[i].append(preds[i])
write_out(OUT_CSV, ids, preds, expl)
no_block = sum(1 for t2 in texts if not re.search(r"answers?\s*:", t2 or "", re.I))
log(f"Pass 1 done in {c1:.0f}s ({no_block}/{len(texts)} had no ANSWERS: block)")
reserve = min(300.0, 0.25 * c1 + 60) if WANT_EXPL else 30.0
ne = 0
while left() - reserve > c1 * 1.25 and ne < VOTE_CAP:
ne += 1; log(f"vote pass {ne} ({left():.0f}s left)")
texts = gen(prompts, max_new, sample=True)
for i, n, txt in zip(ids, ns, texts):
if txt: samples[i].append(parse(txt, n, tasks[i], srcs[i]))
for i, n in zip(ids, ns):
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)],
g[k] if k < len(g) else None) for k in range(n)], n, srcs[i])
write_out(OUT_CSV, ids, preds, expl); log(f"voted over {ne+1} samples")
log(f"self-consistency: {ne} vote pass(es)")
if WANT_EXPL and left() > 60:
ep = [tok.apply_chat_template([{"role": "system", "content": EXPL_SYS},
{"role": "user", "content": f"{str(r['context']).strip()}\n\n{str(r['query']).strip()}\n\n"
f"Answers:\n" + "\n".join(f"- {a}" for a in preds[str(r['id'])]) + "\n\nExplain the rules briefly."}],
tokenize=False, add_generation_prompt=True) for _, r in df.iterrows()]
for i, e in zip(ids, gen(ep, 200, sample=False)):
e = re.sub(r"\s+", " ", (e or "").strip())
if e: expl[i] = e[:1200]
write_out(OUT_CSV, ids, preds, expl); log("explanations written")
for i, n in zip(ids, ns):
if len(preds[i]) != n or any(not str(x).strip() for x in preds[i]):
preds[i] = fit_to_n([x for x in preds[i] if str(x).strip()], n, srcs[i])
write_out(OUT_CSV, ids, preds, expl)
log(f"DONE. {len(ids)} rows, {time.time()-T0:.0f}s.")
if __name__ == "__main__":
main()
|