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
|