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()