File size: 7,475 Bytes
8129b09 22b973c 8129b09 22b973c 8129b09 22b973c 8129b09 1ec7af8 8129b09 1ec7af8 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 | #!/usr/bin/env python3
"""Solving pipeline used by the probe submission (vLLM or transformers stack).
Strategy: budgeted thinking pass -> JSON answers; incremental writes;
stage-0 fast pass first on the transformers path so nothing is ever blank.
"""
import csv
import json
import time
from iol_common import direct_prompt, infer_labels, parse_items, strip_think
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_items": len(labels),
})
return probs
def _write(out_csv, results, order):
with open(out_csv, "w", newline="", encoding="utf-8") as f:
w = csv.DictWriter(f, fieldnames=["id", "pred", "explanation"])
w.writeheader()
for pid in order:
r = results[pid]
w.writerow({"id": pid, "pred": json.dumps(r["pred"], ensure_ascii=False),
"explanation": r.get("explanation", "")})
def _explanation(sample_text):
tail = strip_think(sample_text).strip()
return (tail[:500] + "...") if len(tail) > 500 else tail
def solve_all_vllm(llm, rows, out_csv, start, deadline_s, log):
from vllm import SamplingParams
probs = _rows_to_probs(rows)
order = [p["id"] for p in probs]
results = {p["id"]: {"pred": [""] * p["n_items"], "explanation": ""} for p in probs}
_write(out_csv, results, order)
tok = llm.get_tokenizer()
remaining = deadline_s - (time.time() - start)
# scale the thinking budget to the time left and the row count, assuming
# a conservative ~90 tok/s aggregate on the T4
est_tokens = max(remaining - 120, 60) * 90
think_budget = int(min(3072, max(768, est_tokens / max(len(probs), 1))))
log(f"vllm solve: {len(probs)} rows, think_budget={think_budget}, remaining={remaining:.0f}s")
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)
prompts = [render(direct_prompt(p["context"], p["query"], p["task_type"], p["labels"]), True)
for p in probs]
sp = SamplingParams(temperature=0.0, max_tokens=think_budget + 512)
t = time.time()
outs = llm.generate(prompts, sp)
texts = [o.outputs[0].text for o in outs]
log(f"vllm gen done in {time.time()-t:.0f}s, "
f"{sum(len(o.outputs[0].token_ids) for o in outs)} tokens out")
# budget forcing: close unfinished thinking and squeeze the answer out
todo = [i for i, tx in enumerate(texts) if "</think>" not in tx]
if todo and deadline_s - (time.time() - start) > 90:
log(f"budget-forcing {len(todo)} unclosed samples")
cont_prompts = [prompts[i] + texts[i] +
"\n\nOkay, time is up — I must answer now.\n</think>\n\n"
for i in todo]
cont = llm.generate(cont_prompts, SamplingParams(temperature=0.0, max_tokens=512))
for i, o in zip(todo, cont):
texts[i] = texts[i] + "\n</think>\n\n" + o.outputs[0].text
for p, text in zip(probs, texts):
results[p["id"]] = {"pred": parse_items(text, p["labels"]),
"explanation": _explanation(text)}
_write(out_csv, results, order)
log("vllm solve: submission written")
def solve_all_transformers(model_dir, rows, out_csv, start, deadline_s, log):
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
probs = _rows_to_probs(rows)
order = [p["id"] for p in probs]
results = {p["id"]: {"pred": [""] * p["n_items"], "explanation": ""} for p in probs}
for pid in order[:54]:
results[pid]["explanation"] = "[diag] tf solver started, model loading"
tok = AutoTokenizer.from_pretrained(model_dir)
if tok.pad_token_id is None:
tok.pad_token = tok.eos_token
tok.padding_side = "left"
model = AutoModelForCausalLM.from_pretrained(
model_dir, torch_dtype=torch.float16).eval()
model.to("cuda" if torch.cuda.is_available() else "cpu")
log(f"transformers model loaded from {model_dir}")
for pid in order[:63]:
results[pid]["explanation"] = "[diag] tf model loaded, generating"
_write(out_csv, results, order)
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 gen_batch(prompt_texts, max_new, batch=4):
res = []
for s in range(0, len(prompt_texts), batch):
chunk = prompt_texts[s:s + batch]
enc = tok(chunk, return_tensors="pt", padding=True, truncation=True,
max_length=7000).to(model.device)
with torch.no_grad():
out = model.generate(**enc, max_new_tokens=max_new, do_sample=False,
pad_token_id=tok.pad_token_id)
res.extend(tok.batch_decode(out[:, enc["input_ids"].shape[1]:],
skip_special_tokens=True))
return res
# Stage 0: fast no-think pass, incremental writes so a timeout keeps partials
prompts0 = [render(direct_prompt(p["context"], p["query"], p["task_type"], p["labels"]),
False) for p in probs]
t = time.time()
B = 4
for s in range(0, len(prompts0), B):
if time.time() - start > deadline_s - 60:
log("tf stage0: deadline, stopping")
break
outs0 = gen_batch(prompts0[s:s + B], 400, batch=B)
for p, text in zip(probs[s:s + B], outs0):
results[p["id"]] = {"pred": parse_items(text, p["labels"]),
"explanation": _explanation(text)}
_write(out_csv, results, order)
log(f"tf stage0 {min(s+B, len(probs))}/{len(probs)}")
log(f"tf stage0 done in {time.time()-t:.0f}s")
# Stage 1: thinking pass per problem while time remains
for p in probs:
if time.time() - start > deadline_s - 90:
log("tf stage1: deadline, stopping")
break
try:
text = gen_batch([render(direct_prompt(p["context"], p["query"], p["task_type"],
p["labels"]), True)], 1800, batch=1)[0]
pred = parse_items(text, p["labels"])
if any(x.strip() for x in pred):
old = results[p["id"]]["pred"]
merged = [a if a.strip() else (old[i] if i < len(old) else "")
for i, a in enumerate(pred)]
results[p["id"]] = {"pred": merged, "explanation": _explanation(text)}
_write(out_csv, results, order)
except Exception as e:
log(f"tf stage1 {p['id']}: {e!r}")
_write(out_csv, results, order)
log("tf solve: submission written")
|