File size: 10,903 Bytes
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f15c25d
8129b09
 
 
f15c25d
 
 
 
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f15c25d
 
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a69a5f2
8129b09
005b7d4
a69a5f2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8129b09
 
 
 
 
 
 
 
40b66d7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8129b09
40b66d7
a69a5f2
 
 
40b66d7
005b7d4
 
 
40b66d7
 
005b7d4
40b66d7
 
005b7d4
 
 
a69a5f2
 
005b7d4
a69a5f2
005b7d4
 
a69a5f2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8129b09
 
 
40b66d7
 
8129b09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Production solving pipeline for the IOL-AI 2026 submission (vLLM path).

Stages (all incremental-writing, deadline-aware):
  0. no-think greedy pass            -> guaranteed floor for every row (~1-2 min)
  1. thinking pass, adaptive budget  -> main answers (budget-forced close)
  2. re-solve pass for weak rows     -> rows that were force-closed / have empty
                                        items / are number tasks, at 1.5x budget,
                                        temperature 0.6; per-item weighted vote
"""
import csv
import json
import os
import time

# For rehearsals on faster GPUs: scale the measured token rate to emulate the
# T4 when choosing budgets (e.g. IOL_RATE_SCALE=0.4). 1.0 in production.
RATE_SCALE = float(os.environ.get("IOL_RATE_SCALE", "1.0"))

from iol_common import (direct_prompt, infer_labels, majority_vote, parse_items,
                        strip_think)

TASK_EXPL = {
    "translation": "Worked out vocabulary and word order from the paired examples, then applied them.",
    "fill_blanks": "Worked out the morphological pattern from the example table and applied it to the blanks.",
    "match_letters": "Matched items to meanings by cross-checking recurring morphemes across the examples.",
    "text_to_num": "Derived the number system (base and composition rules) from the examples.",
    "num_to_text": "Derived the number system (base and composition rules) from the examples.",
}


def rows_to_probs(rows):
    probs = []
    for r in rows:
        labels = infer_labels(r.get("context", ""), r.get("query", ""))
        probs.append({
            "id": r["id"], "context": r.get("context", ""),
            "query": r.get("query", ""), "task_type": (r.get("task_type") or "").strip(),
            "labels": labels, "n": len(labels),
        })
    return probs


class Submission:
    def __init__(self, probs, out_csv, diag=None):
        self.out_csv = out_csv
        self.order = [p["id"] for p in probs]
        self.rows = {p["id"]: {"pred": [""] * p["n"], "explanation": ""} for p in probs}
        if diag:
            count, note = diag
            for pid in self.order[:count]:
                self.rows[pid]["explanation"] = f"[diag] {note}"
        self.write()

    def update(self, pid, pred, explanation=None):
        self.rows[pid]["pred"] = pred
        if explanation:
            self.rows[pid]["explanation"] = explanation[:600]

    def write(self):
        with open(self.out_csv, "w", newline="", encoding="utf-8") as f:
            w = csv.DictWriter(f, fieldnames=["id", "pred", "explanation"])
            w.writeheader()
            for pid in self.order:
                r = self.rows[pid]
                w.writerow({"id": pid, "pred": json.dumps(r["pred"], ensure_ascii=False),
                            "explanation": r["explanation"]})


def extract_explanation(text, task_type):
    from iol_common import extract_json
    obj = extract_json(strip_think(text))
    expl = str(obj.get("explanation", "") or "").strip() if obj else ""
    if not expl or expl.startswith("{"):
        expl = TASK_EXPL.get(task_type, TASK_EXPL["translation"])
    return expl


def solve(llm, rows, out_csv, start, deadline_s, log, tok_per_s_guess=90.0):
    from vllm import SamplingParams
    probs = rows_to_probs(rows)
    sub = Submission(probs, out_csv,
                     diag=(45, "vllm engine up, generation starting"))
    tok = llm.get_tokenizer()
    n_rows = max(len(probs), 1)

    def render(p, thinking):
        try:
            return tok.apply_chat_template(
                [{"role": "user", "content": p}], tokenize=False,
                add_generation_prompt=True, enable_thinking=thinking)
        except TypeError:
            return tok.apply_chat_template(
                [{"role": "user", "content": p}], tokenize=False,
                add_generation_prompt=True)

    def shrink(prompt, reserve):
        """Keep the prompt inside the context window (drop middle of context)."""
        ids = tok(prompt)["input_ids"]
        limit = getattr(llm.llm_engine.model_config, "max_model_len", 10240) - reserve
        if len(ids) <= limit:
            return prompt
        keep = limit // 2
        return tok.decode(ids[:keep]) + "\n[...data truncated...]\n" + tok.decode(ids[-keep:])

    def left(): return deadline_s - (time.time() - start)

    base_prompts = {p["id"]: direct_prompt(p["context"], p["query"], p["task_type"],
                                           p["labels"]) for p in probs}
    votes = {p["id"]: [] for p in probs}   # (answers, weight, text)

    # ---------------- stage 0: no-think floor (chunked, incremental) --------
    t0 = time.time()
    sp0 = SamplingParams(temperature=0.0, max_tokens=380)
    stage0_tok = 0
    CHUNK = 24
    for s in range(0, len(probs), CHUNK):
        if left() < 90 and s > 0:
            log("stage0: low on time, stopping early")
            break
        chunk = probs[s:s + CHUNK]
        rend0 = [shrink(render(base_prompts[p["id"]], False), 900) for p in chunk]
        outs = llm.generate(rend0, sp0)
        stage0_tok += sum(len(o.outputs[0].token_ids) for o in outs)
        for p, o in zip(chunk, outs):
            ans = parse_items(o.outputs[0].text, p["labels"])
            votes[p["id"]].append((ans, 1.0, o.outputs[0].text))
            sub.update(p["id"], ans, extract_explanation(o.outputs[0].text, p["task_type"]))
        sub.write()
        log(f"stage0 {min(s+CHUNK, len(probs))}/{len(probs)} written")
    stage0_dt = time.time() - t0
    tok_rate = max(stage0_tok / max(stage0_dt, 1e-6), 20.0) * RATE_SCALE
    log(f"stage0 done in {stage0_dt:.0f}s, {stage0_tok} tok "
        f"(planning rate {tok_rate:.0f} tok/s, scale {RATE_SCALE})")

    # ---------------- stage 1: thinking, adaptive budget ----------------
    if left() < 120:
        return
    # Budget skew, calibrated on held-out IOL 2024: extra thinking pays off on
    # match_letters (EM doubles 1k->4k), is flat on translation (EM 0, chrF
    # unchanged), morpho/number tasks in between. Deep rows go first so a
    # timeout costs the low-ROI rows, which keep their stage0 floor.
    MULT = {"match_letters": 2.5, "fill_blanks": 1.3, "text_to_num": 1.3,
            "num_to_text": 1.3, "translation": 0.6}
    # Order: one cheap low-mult chunk first to calibrate the decode rate
    # (stage0's estimate is prefill-dominated and ~2x too low — an inverted
    # skew starves the high-value rows), then high-mult rows while the budget
    # is both accurate and plentiful, then the rest.
    low = [p for p in probs if MULT.get(p["task_type"], 1.0) <= 1.0]
    high = sorted((p for p in probs if MULT.get(p["task_type"], 1.0) > 1.0),
                  key=lambda p: -MULT.get(p["task_type"], 1.0))
    stage1_order = low[:CHUNK] + high + low[CHUNK:]
    weights = [MULT.get(p["task_type"], 1.0) for p in stage1_order]

    forced = set()
    for s in range(0, len(stage1_order), CHUNK):
        if left() < 120 and s > 0:
            log("stage1: low on time, stopping early")
            break
        chunk = stage1_order[s:s + CHUNK]
        # re-plan the budget per chunk: the rate estimate improves as real
        # decode measurements come in (stage0's rate is prefill-dominated,
        # while stage1 prefill is nearly free thanks to prefix caching)
        weight_left = sum(weights[s:]) or 1.0
        chunk_mult = sum(weights[s:s + CHUNK]) / max(len(chunk), 1)
        budget_tokens = (left() - 90) * tok_rate * 0.9
        think_budget = int(min(6144, max(768,
            (budget_tokens / weight_left) * chunk_mult - 500)))
        log(f"stage1 chunk@{s}: think_budget={think_budget} "
            f"(rate {tok_rate:.0f} tok/s, {left():.0f}s left)")
        sp1 = SamplingParams(temperature=0.0, max_tokens=think_budget + 600)
        rend1 = [shrink(render(base_prompts[p["id"]], True), think_budget + 800)
                 for p in chunk]
        tch = time.time()
        outs = llm.generate(rend1, sp1)
        chunk_tok = sum(len(o.outputs[0].token_ids) for o in outs)
        tok_rate = max(chunk_tok / max(time.time() - tch, 1e-6), 20.0) * RATE_SCALE
        texts = [o.outputs[0].text for o in outs]
        todo = [i for i, tx in enumerate(texts) if "</think>" not in tx]
        if todo and left() > 90:
            cont = llm.generate(
                [rend1[i] + texts[i] + "\n\nOkay, time is up — I must answer now.\n</think>\n\n"
                 for i in todo],
                SamplingParams(temperature=0.0, max_tokens=600))
            for i, o in zip(todo, cont):
                texts[i] += "\n</think>\n\n" + o.outputs[0].text
                forced.add(chunk[i]["id"])
        for p, tx in zip(chunk, texts):
            ans = parse_items(tx, p["labels"])
            w = 2.0 if p["id"] not in forced else 1.5
            votes[p["id"]].append((ans, w, tx))
            sub.update(p["id"], _merge(votes[p["id"]], p["n"]),
                       extract_explanation(tx, p["task_type"]))
        sub.write()
        log(f"stage1 {min(s+CHUNK, len(probs))}/{len(probs)} written ({left():.0f}s left)")

    # ---------------- stage 2: re-solve weak rows ----------------
    weak = [p for p in probs
            if any(not a.strip() for a in sub.rows[p["id"]]["pred"])
            or p["task_type"] in ("text_to_num", "num_to_text", "match_letters")]
    if not weak or left() < 150:
        return
    budget2 = int(min(4096, max(1280, (left() - 150) * tok_rate * 0.7 / len(weak) - 600)))
    log(f"stage2: {len(weak)} weak rows, budget={budget2}")
    sp2 = SamplingParams(temperature=0.6, top_p=0.95, max_tokens=budget2 + 600, seed=1234)
    rend2 = [shrink(render(base_prompts[p["id"]], True), budget2 + 800) for p in weak]
    outs = llm.generate(rend2, sp2)
    texts = [o.outputs[0].text for o in outs]
    todo = [i for i, tx in enumerate(texts) if "</think>" not in tx]
    if todo and left() > 60:
        cont = llm.generate(
            [rend2[i] + texts[i] + "\n\nOkay, time is up — I must answer now.\n</think>\n\n"
             for i in todo],
            SamplingParams(temperature=0.0, max_tokens=600))
        for i, o in zip(todo, cont):
            texts[i] += "\n</think>\n\n" + o.outputs[0].text
    for p, tx in zip(weak, texts):
        ans = parse_items(tx, p["labels"])
        votes[p["id"]].append((ans, 1.5, tx))
        merged = _merge(votes[p["id"]], p["n"])
        sub.update(p["id"], merged)
    sub.write()
    log(f"stage2 written ({left():.0f}s left)")


def _merge(vote_list, n):
    """Per-item weighted vote across passes; never returns empty if any pass answered."""
    out = []
    for j in range(n):
        cands = []
        for ans, w, _ in vote_list:
            if j < len(ans) and ans[j].strip():
                cands.extend([ans[j]] * max(int(w * 2), 1))
        out.append(majority_vote(cands) if cands else "")
    return out